"""Tableau de bord : agrégations par période, comparaisons, séries (règles documentées dans DECISIONS D-04/D-05)."""
from __future__ import annotations

from dataclasses import dataclass
from datetime import date, timedelta
from typing import Any, Literal

from sqlalchemy import func, select, true
from sqlalchemy.orm import Session

from ..core import isoweek
from ..models import Bilan, BilanVersion
from . import bilans as bilan_service

PeriodKind = Literal["week", "month", "year", "all", "custom"]
Rule = Literal["jeudi", "prorata"]


@dataclass
class Period:
    kind: str
    lo: date
    hi: date
    label: str
    short: str
    params: dict[str, Any]


class StatsError(ValueError):
    pass


def make_period(kind: str, *, year: int | None = None, week: int | None = None, month: int | None = None,
                date_from: date | None = None, date_to: date | None = None, bounds: tuple[date, date] | None = None
                ) -> Period:
    if kind == "week":
        assert year and week
        lo, hi = isoweek.week_bounds(year, week)
        return Period(kind, lo, hi, f"Semaine {week} — {isoweek.fr_date(lo)} → {isoweek.fr_date(hi)}",
                      f"S{week} {year}", {"year": year, "week": week})
    if kind == "month":
        assert year and month
        lo, hi = isoweek.month_bounds(year, month)
        return Period(kind, lo, hi, f"{isoweek.MOIS[month - 1].capitalize()} {year}",
                      f"{isoweek.MOIS[month - 1][:4]}. {year}", {"year": year, "month": month})
    if kind == "year":
        assert year
        return Period(kind, date(year, 1, 1), date(year, 12, 31), f"Année {year}", str(year), {"year": year})
    if kind == "custom":
        if not date_from or not date_to or date_from > date_to:
            raise StatsError("Indiquez une date de début antérieure à la date de fin.")
        if (date_to - date_from).days > 366 * 15:
            raise StatsError("Période trop longue (15 ans maximum).")
        return Period(kind, date_from, date_to, f"Du {isoweek.fr_date(date_from)} au {isoweek.fr_date(date_to)}",
                      f"{isoweek.fr_date(date_from)} → {isoweek.fr_date(date_to)}",
                      {"from": date_from.isoformat(), "to": date_to.isoformat()})
    if kind == "all":
        assert bounds
        return Period(kind, bounds[0], bounds[1], "Depuis le début des données", "Depuis toujours", {})
    raise StatsError("Type de période inconnu.")


def shift_period(p: Period, mode: str) -> Period | None:
    if mode == "none" or p.kind == "all":
        return None
    if p.kind == "week":
        y, w = p.params["year"], p.params["week"]
        if mode == "previous":
            y, w = isoweek.shift_week(y, w, -1)
        else:
            y = y - 1
            w = min(w, isoweek.weeks_in_year(y))
        return make_period("week", year=y, week=w)
    if p.kind == "month":
        y, m = p.params["year"], p.params["month"]
        if mode == "previous":
            y, m = (y - 1, 12) if m == 1 else (y, m - 1)
        else:
            y -= 1
        return make_period("month", year=y, month=m)
    if p.kind == "year":
        return make_period("year", year=p.params["year"] - 1)
    length = (p.hi - p.lo).days + 1
    if mode == "previous":
        return make_period("custom", date_from=p.lo - timedelta(days=length), date_to=p.lo - timedelta(days=1))
    try:
        return make_period("custom", date_from=p.lo.replace(year=p.lo.year - 1), date_to=p.hi.replace(year=p.hi.year - 1))
    except ValueError:  # 29 février
        return make_period("custom", date_from=p.lo - timedelta(days=365), date_to=p.hi - timedelta(days=365))


def _weight(start: date, end: date, p: Period, rule: str) -> float:
    if p.kind in {"week", "all"}:
        return 1.0 if start <= p.hi and end >= p.lo else 0.0
    if rule == "prorata":
        return isoweek.days_in_range(start, end, p.lo, p.hi) / 7
    thursday = start + timedelta(days=3)
    return 1.0 if p.lo <= thursday <= p.hi else 0.0


def _official_data(db: Session, lo: date, hi: date, include_drafts: bool) -> list[dict[str, Any]]:
    """Bilans touchant [lo, hi] avec leurs lignes officielles (dernière version validée)."""
    bilans = list(db.execute(select(Bilan).where(Bilan.period_end >= lo, Bilan.period_start <= hi)
                             .order_by(Bilan.period_start)).scalars())
    if not bilans:
        return []
    latest = {v.bilan_id: v for v in db.execute(
        select(BilanVersion).distinct(BilanVersion.bilan_id)
        .where(BilanVersion.bilan_id.in_([b.id for b in bilans]))
        .order_by(BilanVersion.bilan_id, BilanVersion.version_no.desc())).scalars()}
    out = []
    indicators = bilan_service.get_indicators(db) if include_drafts else []
    for b in bilans:
        v = latest.get(b.id)
        draft_state = b.status in bilan_service.EDITABLE
        if include_drafts and draft_state:
            rows, official = bilan_service.rows_for(db, b, indicators), False
        elif v is not None:
            rows, official = v.snapshot["rows"], True
        else:
            continue
        out.append({"year": b.iso_year, "week": b.iso_week, "start": b.period_start, "end": b.period_end,
                    "official": official, "status": b.status,
                    "values": {r["indicator_id"]: r["value"] for r in rows},
                    "updated": (b.validated_at or b.updated_at)})
    return out


def _aggregate(records: list[dict[str, Any]], p: Period, rule: str, ind_ids: list[int]) -> dict[str, Any]:
    sums: dict[int, float] = {i: 0.0 for i in ind_ids}
    filled: dict[int, int] = {i: 0 for i in ind_ids}
    used = []
    for rec in records:
        w = _weight(rec["start"], rec["end"], p, rule)
        if w <= 0:
            continue
        used.append((rec, w))
        for i in ind_ids:
            v = rec["values"].get(i)
            if v is not None:
                sums[i] += v * w
                filled[i] += 1
    values = {i: (round(sums[i], 1) if rule == "prorata" and p.kind not in {"week", "all"} else int(round(sums[i])))
              if filled[i] else None for i in ind_ids}
    return {"values": values, "filled": filled, "records": used}


def _expected_weeks(p: Period, rule: str) -> list[tuple[int, int]]:
    weeks = isoweek.weeks_touching(p.lo, p.hi)
    if p.kind in {"week", "all"} or rule == "prorata":
        return weeks
    return [(y, w) for (y, w) in weeks if p.lo <= isoweek.week_thursday(y, w) <= p.hi]


def _series(db: Session, p: Period, rule: str, include_drafts: bool, ind_ids: list[int]) -> dict[str, Any]:
    """Série d'évolution : hebdomadaire (≤ 16 semaines, ou 12 dernières semaines en vue semaine), sinon mensuelle."""
    if p.kind == "week":
        y, w = p.params["year"], p.params["week"]
        weeks = [isoweek.shift_week(y, w, -k) for k in range(11, -1, -1)]
        lo = isoweek.week_bounds(*weeks[0])[0]
        recs = {(r["year"], r["week"]): r for r in _official_data(db, lo, p.hi, include_drafts)}
        points = [{"label": f"S{wk}", "title": isoweek.week_label(yr, wk), "start": isoweek.week_bounds(yr, wk)[0].isoformat(),
                   "values": {i: (recs[(yr, wk)]["values"].get(i) if (yr, wk) in recs else None) for i in ind_ids},
                   "missing": (yr, wk) not in recs} for yr, wk in weeks]
        return {"granularity": "semaine", "context": "12 dernières semaines", "points": points}
    weeks = _expected_weeks(p, rule)
    records = _official_data(db, p.lo - timedelta(days=6), p.hi + timedelta(days=6), include_drafts)
    if len(weeks) <= 16:
        by_week = {(r["year"], r["week"]): r for r in records}
        points = [{"label": f"S{wk}", "title": isoweek.week_label(yr, wk), "start": isoweek.week_bounds(yr, wk)[0].isoformat(),
                   "values": {i: (by_week[(yr, wk)]["values"].get(i) if (yr, wk) in by_week else None) for i in ind_ids},
                   "missing": (yr, wk) not in by_week} for yr, wk in weeks]
        return {"granularity": "semaine", "context": None, "points": points}
    points = []
    cur = date(p.lo.year, p.lo.month, 1)
    while cur <= p.hi:
        mp = make_period("month", year=cur.year, month=cur.month)
        mp = Period("month", max(mp.lo, p.lo), min(mp.hi, p.hi), mp.label, mp.short, mp.params)
        agg = _aggregate(records, mp, rule, ind_ids)
        points.append({"label": f"{isoweek.MOIS[cur.month - 1][:3]}. {str(cur.year)[2:]}", "title": mp.label,
                       "start": cur.isoformat(), "values": agg["values"], "missing": not agg["records"]})
        cur = date(cur.year + (cur.month == 12), cur.month % 12 + 1, 1)
    return {"granularity": "mois", "context": None, "points": points}


def compute(db: Session, kind: str, *, year: int | None = None, week: int | None = None, month: int | None = None,
            date_from: date | None = None, date_to: date | None = None, compare: str = "previous",
            rule: str = "jeudi", include_drafts: bool = False) -> dict[str, Any]:
    first = db.execute(select(func.min(Bilan.period_start), func.max(Bilan.period_end), func.count(Bilan.id))
                       .where(Bilan.current_version > 0 if not include_drafts else true())).one()
    if kind == "all":
        if first[0] is None:
            bounds = (date.today(), date.today())
        else:
            bounds = (first[0], first[1])
        p = make_period("all", bounds=bounds)
    else:
        p = make_period(kind, year=year, week=week, month=month, date_from=date_from, date_to=date_to)
    indicators = bilan_service.get_indicators(db)
    ind_ids = [i.id for i in indicators]
    records = _official_data(db, p.lo - timedelta(days=6), p.hi + timedelta(days=6), include_drafts)
    agg = _aggregate(records, p, rule, ind_ids)
    cmp_p = shift_period(p, compare)
    cmp_agg = None
    if cmp_p is not None:
        cmp_records = _official_data(db, cmp_p.lo - timedelta(days=6), cmp_p.hi + timedelta(days=6), include_drafts)
        cmp_agg = _aggregate(cmp_records, cmp_p, rule, ind_ids)
    expected = _expected_weeks(p, rule)
    used_weeks = {(r["year"], r["week"]) for r, _ in agg["records"]}
    straddling = []
    if p.kind in {"month", "year", "custom"}:
        for (y, w) in isoweek.weeks_touching(p.lo, p.hi):
            s, e = isoweek.week_bounds(y, w)
            if s < p.lo or e > p.hi:
                inside = isoweek.days_in_range(s, e, p.lo, p.hi)
                th = isoweek.week_thursday(y, w)
                if rule == "prorata":
                    note = f"comptée pour {inside}/7 de ses valeurs (prorata)"
                else:
                    note = ("comptée dans cette période (jeudi " + isoweek.fr_date(th) + ")" if p.lo <= th <= p.hi
                            else "exclue de cette période (jeudi " + isoweek.fr_date(th) + ")")
                straddling.append({"year": y, "week": w, "start": s.isoformat(), "end": e.isoformat(),
                                   "days_inside": inside, "note": note})
    series = _series(db, p, rule, include_drafts, ind_ids)
    rows = []
    total = total_cmp = 0.0
    total_has = total_cmp_has = False
    for ind in indicators:
        v = agg["values"][ind.id]
        if not ind.is_active and v is None:
            continue
        c = cmp_agg["values"][ind.id] if cmp_agg else None
        delta = pct = None
        tone = "neutre"
        if v is not None and c is not None:
            delta = round(v - c, 1)
            pct = None if c == 0 else round(100 * (v - c) / c, 1)
            if delta != 0 and ind.direction:
                good = (delta > 0) == (ind.direction == "hausse_favorable")
                tone = "favorable" if good else "defavorable"
        if ind.include_in_total:
            if v is not None:
                total += v
                total_has = True
            if c is not None:
                total_cmp += c
                total_cmp_has = True
        rows.append({"indicator_id": ind.id, "code": ind.code, "label": ind.label, "is_active": ind.is_active,
                     "direction": ind.direction, "include_in_total": ind.include_in_total, "value": v,
                     "compare_value": c, "delta": delta, "pct": pct, "tone": tone,
                     "filled_weeks": agg["filled"][ind.id], "series": [pt["values"].get(ind.id) for pt in series["points"]]})
    last_update = max((r["updated"] for r, _ in agg["records"] if r["updated"]), default=None)
    drafts_used = sum(1 for r, _ in agg["records"] if not r["official"])
    return {
        "period": {"kind": p.kind, "label": p.label, "short": p.short, "start": p.lo.isoformat(),
                   "end": p.hi.isoformat(), **p.params},
        "compare": None if cmp_p is None else {"mode": compare, "label": cmp_p.label, "short": cmp_p.short,
                                                "start": cmp_p.lo.isoformat(), "end": cmp_p.hi.isoformat(),
                                                "bilans": len(cmp_agg["records"]) if cmp_agg else 0},
        "rule": rule, "include_drafts": include_drafts, "drafts_used": drafts_used,
        "rows": rows,
        "total": round(total, 1) if total_has else None,
        "total_compare": round(total_cmp, 1) if total_cmp_has else None,
        "series": {"granularity": series["granularity"], "context": series["context"],
                   "points": [{k: pt[k] for k in ("label", "title", "start", "missing")} for pt in series["points"]]},
        "coverage": {"expected_weeks": len(expected), "weeks_with_data": len(used_weeks & set(expected)) if expected else 0,
                     "missing_weeks": [f"S{w} {y}" for (y, w) in expected if (y, w) not in used_weeks][:60]},
        "straddling": straddling,
        "bilans_count": len(agg["records"]),
        "data_range": {"first": first[0].isoformat() if first[0] else None,
                       "last": first[1].isoformat() if first[1] else None, "bilans": first[2]},
        "last_update": last_update.isoformat() if last_update else None,
    }
