"""Small SQL reporting primitives. No operational writes or row locks."""
from dataclasses import dataclass

from django.db import connections
from django.db.models import Aggregate, Avg, Count, DurationField, ExpressionWrapper, F, Value
from django.db.models.functions import TruncDate


class Percentile(Aggregate):
    function = "PERCENTILE_CONT"
    template = "%(function)s(%(percentile)s) WITHIN GROUP (ORDER BY %(expressions)s)"
    output_field = DurationField()

    def __init__(self, expression, percentile):
        if percentile not in (0.5, 0.9):
            raise ValueError("Only median and P90 are supported.")
        super().__init__(expression, percentile=percentile)


def durations(rows, start, end):
    rows = rows.filter(**{start + "__isnull": False, end + "__isnull": False, end + "__gte": F(start)})
    rows = rows.annotate(report_duration=ExpressionWrapper(F(end) - F(start), output_field=DurationField()))
    return rows.aggregate(count=Count("pk"), average=Avg("report_duration"), median=Percentile("report_duration", 0.5), p90=Percentile("report_duration", 0.9))


def grouped(rows, dimensions, **metrics):
    return rows.order_by().values(*dimensions).annotate(**(metrics or {"count": Count("pk", distinct=True)})).order_by(*dimensions)


def daily(rows, timestamp, **metrics):
    return grouped(rows.annotate(day=TruncDate(timestamp)), ["day"], **metrics)


def totals(rows, **metrics):
    """Lazy form of ``rows.aggregate(**metrics)``: one row, no GROUP BY, combinable.

    Unlike ``aggregate()``, a metric name may not reuse a model field name.
    """
    return rows.order_by().annotate(_whole=Value(1)).values("_whole").annotate(**metrics).values(*metrics)


def combined_totals(**reads):
    """Evaluate several ``totals()`` querysets in one round trip; each equals ``.aggregate()``."""
    querysets = list(reads.values())
    # An aggregate without GROUP BY yields exactly one row, so the cross join is 1 x 1 x ...
    db = querysets[0].db
    sources, params = [], []
    for index, rows in enumerate(querysets):
        sql, values = rows.query.get_compiler(using=db).as_sql()
        sources.append(f"({sql}) AS r{index}")
        params.extend(values)
    with connections[db].cursor() as cursor:
        cursor.execute("SELECT * FROM " + " CROSS JOIN ".join(sources), params)
        row, names = cursor.fetchone(), [column.name for column in cursor.description]
    result, start = {}, 0
    for name, rows in reads.items():
        width = len(rows.query.annotation_select)
        result[name] = dict(zip(names[start:start + width], row[start:start + width]))
        start += width
    return result


def combined_exists(**reads):
    """Evaluate several ``queryset.exists()`` probes in one round trip, with identical SQL."""
    querysets = list(reads.values())
    db = querysets[0].db
    probes, params = [], []
    for rows in querysets:
        sql, values = rows.query.exists().get_compiler(using=db).as_sql()
        probes.append(f"EXISTS({sql})")
        params.extend(values)
    with connections[db].cursor() as cursor:
        cursor.execute("SELECT " + ", ".join(probes), params)
        return dict(zip(reads, cursor.fetchone()))


@dataclass
class Table:
    key: str
    title: str
    rows: object
    columns: tuple
    note: str = ""

    @property
    def headers(self):
        return [label for _, label in self.columns]

    def values(self, rows):
        for row in rows:
            yield [row.get(field) for field, _ in self.columns]


def table(key, title, rows, fields, note=""):
    return Table(key, title, rows, tuple((field, label) for field, label in fields), note)
