"""Logique métier des bilans hebdomadaires : brouillon, contrôles, validation, versions."""
from __future__ import annotations

import hashlib
import json
import uuid
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from decimal import Decimal, InvalidOperation
from typing import Any

from sqlalchemy import select
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.orm import Session, selectinload

from ..core import audit, isoweek
from ..models import (
    STATUS_LABELS, Bilan, BilanValue, BilanVersion, ImportApplication, ImportFile, Indicator, User, ValueDetail,
    ValueHistory,
)
from . import settings as settings_service

MAX_VALUE = 100_000
MAX_OBSERVATION = 2000
MAX_DETAILS = 60
EDITABLE = {"brouillon", "a_verifier", "rouvert"}
COALESCE_WINDOW = timedelta(minutes=5)


class BilanError(Exception):
    def __init__(self, status: int, message: str, code: str = "bilan_error", extra: dict[str, Any] | None = None):
        super().__init__(message)
        self.status, self.message, self.code, self.extra = status, message, code, extra or {}


def utcnow() -> datetime:
    return datetime.now(timezone.utc)


# --- Validation des saisies ----------------------------------------------------------------------

@dataclass
class RowInput:
    indicator_id: int
    value: int | None
    observation: str
    details: list[Decimal] | None = None  # None = ne pas toucher aux détails


@dataclass
class CellIssue:
    indicator_id: int
    field: str
    message: str
    blocking: bool = True


def parse_count(raw: Any) -> tuple[int | None, str | None]:
    """Interprète une valeur « Nombre ». Vide → None (non renseigné), jamais 0."""
    if raw is None:
        return None, None
    if isinstance(raw, bool):
        return None, "Valeur invalide."
    if isinstance(raw, int):
        value = raw
    elif isinstance(raw, float):
        if raw != int(raw):
            return None, "Le nombre doit être un entier (pas de décimale)."
        value = int(raw)
    else:
        text = str(raw).strip().replace(" ", "").replace(" ", "").replace(" ", "")
        if text == "":
            return None, None
        if text.lstrip("-+").replace(",", "").replace(".", "").isdigit() and ("," in text or "." in text):
            try:
                dec = Decimal(text.replace(",", "."))
            except InvalidOperation:
                return None, "Valeur numérique illisible."
            if dec != dec.to_integral_value():
                return None, "Le nombre doit être un entier (pas de décimale)."
            return None, "Le nombre doit être un entier : retirez la partie décimale."
        if not text.lstrip("-+").isdigit():
            return None, "Seuls des chiffres sont acceptés dans la colonne « Nombre »."
        value = int(text)
    if value < 0:
        return None, "Une valeur négative est interdite."
    if value > MAX_VALUE:
        return None, f"Valeur trop grande (maximum {MAX_VALUE:,})".replace(",", " ") + "."
    return value, None


def parse_detail(raw: Any) -> Decimal:
    text = str(raw).strip().replace(",", ".").replace(" ", "").replace(" ", "")
    try:
        dec = Decimal(text)
    except InvalidOperation:
        raise ValueError(f"« {str(raw)[:20]} » n'est pas un nombre.") from None
    if dec < 0 or dec > 100000 or dec != dec.quantize(Decimal("0.001")):
        raise ValueError(f"« {str(raw)[:20]} » : nombre positif avec 3 décimales maximum attendu.")
    return dec


def clean_text(text: Any, limit: int = MAX_OBSERVATION) -> str:
    s = "" if text is None else str(text)
    s = "".join(ch for ch in s if ch in "\n\t" or ch >= " ").replace("\r\n", "\n").replace("\r", "\n")
    return s.strip()[:limit]


# --- Lecture -------------------------------------------------------------------------------------

def get_indicators(db: Session, include_inactive: bool = True) -> list[Indicator]:
    stmt = select(Indicator).order_by(Indicator.position, Indicator.id)
    if not include_inactive:
        stmt = stmt.where(Indicator.is_active.is_(True))
    return list(db.execute(stmt).scalars())


def find_bilan(db: Session, iso_year: int, iso_week: int, *, lock: bool = False) -> Bilan | None:
    stmt = select(Bilan).where(Bilan.iso_year == iso_year, Bilan.iso_week == iso_week)
    if lock:
        stmt = stmt.with_for_update()
    else:
        stmt = stmt.options(selectinload(Bilan.values).selectinload(BilanValue.details))
    return db.execute(stmt).scalar_one_or_none()


def _user_names(db: Session, ids: set[int | None]) -> dict[int, str]:
    ids_clean = {i for i in ids if i}
    if not ids_clean:
        return {}
    return {u.id: u.display_name for u in db.execute(select(User).where(User.id.in_(ids_clean))).scalars()}


def _import_names(db: Session, ids: set[uuid.UUID | None]) -> dict[uuid.UUID, str]:
    ids_clean = {i for i in ids if i}
    if not ids_clean:
        return {}
    return {f.id: f.original_name for f in db.execute(select(ImportFile).where(ImportFile.id.in_(ids_clean))).scalars()}


def _detail_dict(d: ValueDetail) -> dict[str, Any]:
    return {"kind": d.kind, "value": format(d.amount.normalize(), "f"), "unit": d.unit}


def rows_for(db: Session, bilan: Bilan | None, indicators: list[Indicator]) -> list[dict[str, Any]]:
    """Lignes affichées : indicateurs actifs + indicateurs inactifs ayant une valeur dans ce bilan."""
    values = {v.indicator_id: v for v in (bilan.values if bilan else [])}
    imports = _import_names(db, {v.import_id for v in values.values()})
    rows = []
    for ind in indicators:
        v = values.get(ind.id)
        if v is None and not ind.is_active:
            continue
        if v is None and bilan is not None and bilan.status in {"valide", "archive"}:
            continue
        rows.append({
            "indicator_id": ind.id, "code": ind.code,
            "label": v.label_snapshot if v is not None else ind.label,
            "current_label": ind.label, "position": ind.position, "is_active": ind.is_active,
            "detail_kind": ind.detail_kind, "detail_unit": ind.detail_unit, "warn_max": ind.warn_max,
            "value": v.value if v else None,
            "observation": v.observation if v else "",
            "details": [_detail_dict(d) for d in v.details] if v else [],
            "source": None if v is None else {
                "kind": v.source_kind, "import_id": str(v.import_id) if v.import_id else None,
                "file_name": imports.get(v.import_id) if v.import_id else None,
                "locator": v.source_locator, "ref": v.source_ref, "bbox": v.source_bbox,
                "confidence": v.confidence,
            },
            "updated_at": v.updated_at.isoformat() if v and v.updated_at else None,
        })
    return rows


def bilan_payload(db: Session, iso_year: int, iso_week: int, bilan: Bilan | None = None) -> dict[str, Any]:
    start, end = isoweek.week_bounds(iso_year, iso_week)
    if bilan is None:
        bilan = find_bilan(db, iso_year, iso_week)
    indicators = get_indicators(db)
    status = bilan.status if bilan else "nouveau"
    names = _user_names(db, {bilan.validated_by, bilan.updated_by, bilan.reopened_by} if bilan else set())
    cy, cw = isoweek.current_week()
    return {
        "iso_year": iso_year, "iso_week": iso_week,
        "period_start": start.isoformat(), "period_end": end.isoformat(),
        "label": isoweek.week_label(iso_year, iso_week),
        "long_label": f"du {isoweek.fr_long_date(start)} au {isoweek.fr_long_date(end)}",
        "is_current_week": (iso_year, iso_week) == (cy, cw),
        "weeks_in_year": isoweek.weeks_in_year(iso_year),
        "exists": bilan is not None, "status": status, "status_label": STATUS_LABELS[status],
        "editable": bilan is None or bilan.status in EDITABLE,
        "revision": bilan.revision if bilan else 0,
        "current_version": bilan.current_version if bilan else 0,
        "is_demo": bool(bilan and bilan.is_demo),
        "updated_at": bilan.updated_at.isoformat() if bilan else None,
        "updated_by": names.get(bilan.updated_by) if bilan and bilan.updated_by else None,
        "validated_at": bilan.validated_at.isoformat() if bilan and bilan.validated_at else None,
        "validated_by": names.get(bilan.validated_by) if bilan and bilan.validated_by else None,
        "reopened_at": bilan.reopened_at.isoformat() if bilan and bilan.reopened_at else None,
        "reopened_by": names.get(bilan.reopened_by) if bilan and bilan.reopened_by else None,
        "reopen_reason": bilan.reopen_reason if bilan else None,
        "rows": rows_for(db, bilan, indicators),
    }


def latest_version(db: Session, bilan_id: int) -> BilanVersion | None:
    return db.execute(select(BilanVersion).where(BilanVersion.bilan_id == bilan_id)
                      .order_by(BilanVersion.version_no.desc()).limit(1)).scalar_one_or_none()


def snapshot_of(db: Session, bilan: Bilan) -> dict[str, Any]:
    payload = bilan_payload(db, bilan.iso_year, bilan.iso_week, bilan)
    keep = ("iso_year", "iso_week", "period_start", "period_end", "label")
    snap = {k: payload[k] for k in keep}
    snap["rows"] = [
        {k: r[k] for k in ("indicator_id", "code", "label", "position", "detail_kind", "detail_unit", "value",
                           "observation", "details", "source")}
        for r in payload["rows"]
    ]
    return snap


def snapshot_digest(snapshot: dict[str, Any]) -> str:
    return hashlib.sha256(json.dumps(snapshot, sort_keys=True, ensure_ascii=False,
                                     separators=(",", ":"), default=str).encode()).hexdigest()


def official_rows(db: Session, bilan: Bilan, include_drafts: bool = False) -> list[dict[str, Any]] | None:
    """Données officielles = dernière version validée ; brouillon seulement sur demande explicite."""
    if bilan.current_version > 0:
        version = latest_version(db, bilan.id)
        if version is not None:
            return version.snapshot["rows"]
    if include_drafts:
        return rows_for(db, bilan, get_indicators(db))
    return None


# --- Écriture du brouillon -----------------------------------------------------------------------

def get_or_create_for_update(db: Session, iso_year: int, iso_week: int, user_id: int) -> Bilan:
    start, end = isoweek.week_bounds(iso_year, iso_week)
    db.execute(pg_insert(Bilan).values(
        iso_year=iso_year, iso_week=iso_week, period_start=start, period_end=end, status="brouillon",
        revision=0, current_version=0, is_demo=False, created_by=user_id, updated_by=user_id,
    ).on_conflict_do_nothing(constraint="uq_bilans_week"))
    bilan = find_bilan(db, iso_year, iso_week, lock=True)
    assert bilan is not None
    return bilan


def _record_change(db: Session, bilan: Bilan, indicator_id: int, field_name: str, old: str | None, new: str | None,
                   user_id: int, source: str, import_id: uuid.UUID | None, reason: str | None) -> None:
    if old == new:
        return
    now = utcnow()
    if source == "manuel" and bilan.status in {"brouillon", "a_verifier"}:
        # Regroupe les frappes successives (sauvegarde automatique) d'un même utilisateur sur un même champ.
        last = db.execute(select(ValueHistory).where(
            ValueHistory.bilan_id == bilan.id, ValueHistory.indicator_id == indicator_id,
            ValueHistory.field == field_name, ValueHistory.user_id == user_id,
            ValueHistory.source_kind == "manuel",
        ).order_by(ValueHistory.id.desc()).limit(1)).scalar_one_or_none()
        if last is not None and now - last.last_at < COALESCE_WINDOW:
            last.new_value, last.last_at = new, now
            return
    db.add(ValueHistory(bilan_id=bilan.id, indicator_id=indicator_id, field=field_name, old_value=old,
                        new_value=new, source_kind=source, import_id=import_id, reason=reason,
                        user_id=user_id, at=now, last_at=now))


def _fmt(v: int | None) -> str | None:
    return None if v is None else str(v)


def _details_str(details: list[ValueDetail] | list[dict[str, Any]]) -> str:
    out = []
    for d in details:
        out.append(format(d.amount.normalize(), "f") if isinstance(d, ValueDetail) else str(d["value"]))
    return " ; ".join(out)


def apply_rows(db: Session, bilan: Bilan, rows: list[RowInput], user_id: int, *, source: str = "manuel",
               import_id: uuid.UUID | None = None, provenance: dict[int, dict[str, Any]] | None = None,
               reason: str | None = None) -> int:
    """Écrit des lignes dans un bilan modifiable. Renvoie le nombre de lignes réellement modifiées."""
    indicators = {i.id: i for i in get_indicators(db)}
    existing = {v.indicator_id: v for v in db.execute(
        select(BilanValue).where(BilanValue.bilan_id == bilan.id).options(selectinload(BilanValue.details))
    ).scalars()}
    changed = 0
    now = utcnow()
    for row in rows:
        ind = indicators.get(row.indicator_id)
        if ind is None:
            raise BilanError(422, "Indicateur inconnu.", "unknown_indicator")
        v = existing.get(row.indicator_id)
        if v is None:
            if not ind.is_active:
                raise BilanError(422, f"L'indicateur « {ind.label} » est désactivé.", "inactive_indicator")
            if row.value is None and not row.observation and not row.details:
                continue
            v = BilanValue(bilan_id=bilan.id, indicator_id=ind.id, label_snapshot=ind.label, value=None,
                           observation="", source_kind=source, created_by=user_id, warnings=[])
            db.add(v)
            db.flush()
            existing[ind.id] = v
            v.details = []
        row_changed = False
        if v.value != row.value:
            _record_change(db, bilan, ind.id, "valeur", _fmt(v.value), _fmt(row.value), user_id, source, import_id,
                           reason)
            v.value, row_changed = row.value, True
        if v.observation != row.observation:
            _record_change(db, bilan, ind.id, "observation", v.observation or None, row.observation or None,
                           user_id, source, import_id, reason)
            v.observation, row_changed = row.observation, True
        if row.details is not None and ind.detail_kind:
            old = _details_str(v.details)
            new = " ; ".join(format(d.normalize(), "f") for d in row.details)
            if old != new:
                _record_change(db, bilan, ind.id, "details", old or None, new or None, user_id, source, import_id,
                               reason)
                v.details.clear()
                db.flush()
                for pos, amount in enumerate(row.details):
                    v.details.append(ValueDetail(kind=ind.detail_kind, amount=amount, unit=ind.detail_unit,
                                                 position=pos))
                row_changed = True
        if row_changed:
            changed += 1
            v.updated_by, v.updated_at = user_id, now
            v.source_kind, v.import_id = source, import_id
            prov = (provenance or {}).get(ind.id, {})
            v.source_locator = prov.get("locator")
            v.source_ref = prov.get("ref")
            v.source_bbox = prov.get("bbox")
            v.confidence = prov.get("confidence")
            if v.label_snapshot != ind.label and bilan.current_version == 0:
                v.label_snapshot = ind.label
    db.flush()
    db.expire(bilan, ["values"])  # la collection chargée en mémoire ne reflète pas les lignes ajoutées
    return changed


def save_draft(db: Session, iso_year: int, iso_week: int, revision: int, raw_rows: list[dict[str, Any]],
               user: User, *, source: str = "manuel", ip: str | None = None) -> dict[str, Any]:
    issues: list[CellIssue] = []
    rows: list[RowInput] = []
    for r in raw_rows:
        ind_id = int(r.get("indicator_id", 0))
        value, err = parse_count(r.get("value"))
        if err:
            issues.append(CellIssue(ind_id, "value", err))
        details: list[Decimal] | None = None
        if "details" in r and r["details"] is not None:
            details = []
            raw_details = r["details"]
            if not isinstance(raw_details, list) or len(raw_details) > MAX_DETAILS:
                issues.append(CellIssue(ind_id, "details", "Liste de détails invalide."))
                raw_details = []
            for item in raw_details:
                raw = item.get("value") if isinstance(item, dict) else item
                if raw is None or str(raw).strip() == "":
                    continue
                try:
                    details.append(parse_detail(raw))
                except ValueError as e:
                    issues.append(CellIssue(ind_id, "details", str(e)))
        rows.append(RowInput(ind_id, value, clean_text(r.get("observation")), details))
    if issues:
        raise BilanError(422, "Certaines cellules contiennent des erreurs.", "invalid_cells",
                         {"issues": [i.__dict__ for i in issues]})
    bilan = get_or_create_for_update(db, iso_year, iso_week, user.id)
    if bilan.status not in EDITABLE:
        raise BilanError(423, f"Ce bilan est {STATUS_LABELS[bilan.status].lower()} : il doit être rouvert "
                              "avant toute modification.", "locked")
    if bilan.revision != revision:
        raise BilanError(409, "Ce bilan a été modifié entre-temps (autre onglet ou autre utilisateur). "
                              "Rechargez pour récupérer la dernière version.", "conflict",
                         {"current_revision": bilan.revision})
    created = bilan.revision == 0
    changed = apply_rows(db, bilan, rows, user.id, source=source)
    if changed or created:
        bilan.revision += 1
        bilan.updated_by, bilan.updated_at = user.id, utcnow()
    if created:
        audit.record(db, "bilan.created", user_id=user.id, username=user.username, target_type="bilan",
                     target_id=f"{iso_year}-S{iso_week:02d}", ip=ip)
    db.commit()
    return bilan_payload(db, iso_year, iso_week)


# --- Contrôles avant validation ------------------------------------------------------------------

def previous_rows(db: Session, iso_year: int, iso_week: int) -> tuple[tuple[int, int], dict[int, dict[str, Any]]]:
    py, pw = isoweek.shift_week(iso_year, iso_week, -1)
    prev = find_bilan(db, py, pw)
    rows = official_rows(db, prev, include_drafts=True) if prev else None
    return (py, pw), {r["indicator_id"]: r for r in (rows or [])}


def check(db: Session, iso_year: int, iso_week: int) -> dict[str, Any]:
    payload = bilan_payload(db, iso_year, iso_week)
    pct = int(settings_service.get(db, "variation_pct"))
    min_abs = int(settings_service.get(db, "variation_min_abs"))
    (py, pw), prev = previous_rows(db, iso_year, iso_week)
    errors, warnings, empty, zeros, observations, details, diffs = [], [], [], [], [], [], []
    for r in payload["rows"]:
        label = r["label"]
        v = r["value"]
        if v is None:
            empty.append(label)
        elif v == 0:
            zeros.append(label)
        if v is not None and r["warn_max"] is not None and v > r["warn_max"]:
            warnings.append({"indicator_id": r["indicator_id"], "label": label, "kind": "seuil",
                             "message": f"{v} dépasse le seuil d'alerte configuré ({r['warn_max']})."})
        if r["observation"]:
            observations.append({"label": label, "text": r["observation"]})
        if r["details"]:
            details.append({"label": label, "kind": r["detail_kind"], "unit": r["detail_unit"],
                            "values": [d["value"] for d in r["details"]]})
            if v is not None and len(r["details"]) > v:
                warnings.append({"indicator_id": r["indicator_id"], "label": label, "kind": "details",
                                 "message": f"{len(r['details'])} détail(s) saisi(s) pour {v} infraction(s)."})
        p = prev.get(r["indicator_id"], {}).get("value")
        if v is not None and p is not None and v != p:
            delta = v - p
            rel = None if p == 0 else round(100 * delta / p)
            big = abs(delta) >= min_abs and (p == 0 or abs(delta) * 100 >= pct * p)
            diffs.append({"label": label, "previous": p, "current": v, "delta": delta, "pct": rel, "important": big})
            if big:
                warnings.append({"indicator_id": r["indicator_id"], "label": label, "kind": "variation",
                                 "message": f"Écart important avec la semaine {pw} : {p} → {v}"
                                            + (f" ({rel:+d} %)." if rel is not None else ".")})
        if r["source"] and r["source"]["confidence"] is not None and r["source"]["confidence"] < 0.8:
            warnings.append({"indicator_id": r["indicator_id"], "label": label, "kind": "confiance",
                             "message": f"Valeur importée avec une confiance de {round(100 * r['source']['confidence'])} %."})
    duplicates = []
    filled = [r for r in payload["rows"] if r["value"] is not None]
    if filled and prev and all(prev.get(r["indicator_id"], {}).get("value") == r["value"] for r in payload["rows"]
                                if r["indicator_id"] in prev) and len(filled) >= 5:
        duplicates.append(f"Toutes les valeurs sont identiques à la semaine {pw} : vérifiez qu'il ne s'agit pas "
                          "d'un doublon.")
    if payload["exists"]:
        bilan = find_bilan(db, iso_year, iso_week)
        assert bilan is not None
        apps = db.execute(select(ImportApplication, ImportFile).join(ImportFile, ImportFile.id == ImportApplication.import_id)
                          .where(ImportApplication.bilan_id == bilan.id, ImportApplication.reverted_at.is_(None))).all()
        seen: dict[str, str] = {}
        for _app, f in apps:
            if f.sha256 in seen:
                duplicates.append(f"Le fichier « {f.original_name} » a été importé deux fois dans ce bilan.")
            seen[f.sha256] = f.original_name
    if not filled:
        errors.append({"label": None, "message": "Aucune valeur n'est renseignée."})
    return {
        "iso_year": iso_year, "iso_week": iso_week, "label": payload["label"], "long_label": payload["long_label"],
        "status": payload["status"], "status_label": payload["status_label"], "revision": payload["revision"],
        "errors": errors, "warnings": warnings, "empty": empty, "zeros": zeros, "observations": observations,
        "details": details, "previous_week": {"iso_year": py, "iso_week": pw, "available": bool(prev)},
        "diffs": diffs, "duplicates": duplicates,
        "filled": len(filled), "total": len(payload["rows"]),
        "can_validate": not errors and payload["status"] in EDITABLE and payload["exists"],
    }


# --- Cycle de vie --------------------------------------------------------------------------------

def _locked_existing(db: Session, iso_year: int, iso_week: int) -> Bilan:
    bilan = find_bilan(db, iso_year, iso_week, lock=True)
    if bilan is None:
        raise BilanError(404, "Aucun bilan n'existe encore pour cette semaine.", "not_found")
    return bilan


def _week_ref(b: Bilan) -> str:
    return f"{b.iso_year}-S{b.iso_week:02d}"


def validate(db: Session, iso_year: int, iso_week: int, revision: int, user: User, comment: str | None,
             ip: str | None) -> dict[str, Any]:
    summary = check(db, iso_year, iso_week)
    bilan = _locked_existing(db, iso_year, iso_week)
    if bilan.status not in EDITABLE:
        raise BilanError(409, "Ce bilan est déjà validé ou archivé.", "already_validated")
    if bilan.revision != revision:
        raise BilanError(409, "Le bilan a changé depuis la vérification. Relancez la vérification.", "conflict",
                         {"current_revision": bilan.revision})
    if summary["errors"]:
        raise BilanError(422, "Le bilan contient des erreurs bloquantes.", "blocking_errors", {"check": summary})
    snapshot = snapshot_of(db, bilan)
    now = utcnow()
    version_no = bilan.current_version + 1
    snapshot["version_no"] = version_no
    snapshot["validated_by"] = user.display_name
    snapshot["validated_at"] = now.isoformat()
    digest = snapshot_digest(snapshot)
    db.add(BilanVersion(bilan_id=bilan.id, version_no=version_no, snapshot=snapshot, snapshot_sha256=digest,
                        comment=clean_text(comment, 500) or None, is_demo=bilan.is_demo, created_by=user.id))
    previous_status = bilan.status
    bilan.status, bilan.current_version = "valide", version_no
    bilan.validated_by, bilan.validated_at = user.id, now
    bilan.revision += 1
    bilan.updated_by, bilan.updated_at = user.id, now
    audit.record(db, "bilan.validated", user_id=user.id, username=user.username, target_type="bilan",
                 target_id=_week_ref(bilan), ip=ip,
                 details={"version": version_no, "sha256": digest, "statut_precedent": previous_status,
                          "avertissements": len(summary["warnings"]), "lignes_vides": len(summary["empty"])})
    db.commit()
    return bilan_payload(db, iso_year, iso_week)


def reopen(db: Session, iso_year: int, iso_week: int, reason: str, user: User, ip: str | None) -> dict[str, Any]:
    reason = clean_text(reason, 500)
    if len(reason) < 5:
        raise BilanError(422, "Indiquez le motif de la réouverture (5 caractères minimum).", "reason_required")
    bilan = _locked_existing(db, iso_year, iso_week)
    if bilan.status not in {"valide", "archive"}:
        raise BilanError(409, "Seul un bilan validé ou archivé peut être rouvert.", "not_validated")
    previous = bilan.status
    bilan.status = "rouvert"
    bilan.reopened_by, bilan.reopened_at, bilan.reopen_reason = user.id, utcnow(), reason
    bilan.archived_at = None
    bilan.revision += 1
    audit.record(db, "bilan.reopened", user_id=user.id, username=user.username, target_type="bilan",
                 target_id=_week_ref(bilan), ip=ip,
                 details={"motif": reason, "version_conservee": bilan.current_version, "statut_precedent": previous})
    db.commit()
    return bilan_payload(db, iso_year, iso_week)


def set_status(db: Session, iso_year: int, iso_week: int, target: str, user: User, ip: str | None) -> dict[str, Any]:
    bilan = _locked_existing(db, iso_year, iso_week)
    allowed = {
        ("brouillon", "a_verifier"), ("a_verifier", "brouillon"),
        ("valide", "archive"), ("archive", "valide"),
    }
    if (bilan.status, target) not in allowed:
        raise BilanError(409, f"Passage de « {STATUS_LABELS[bilan.status]} » à « {STATUS_LABELS.get(target, target)} »"
                              " impossible.", "bad_transition")
    previous = bilan.status
    bilan.status = target
    bilan.archived_at = utcnow() if target == "archive" else None
    bilan.revision += 1
    action = {"archive": "bilan.archived", "valide": "bilan.unarchived"}.get(target, "bilan.status_changed")
    audit.record(db, action, user_id=user.id, username=user.username, target_type="bilan",
                 target_id=_week_ref(bilan), ip=ip, details={"de": previous, "vers": target})
    db.commit()
    return bilan_payload(db, iso_year, iso_week)


def _restore_rows_from_snapshot(db: Session, bilan: Bilan, snapshot: dict[str, Any], user_id: int, source: str,
                                reason: str | None) -> int:
    known = {i.id for i in get_indicators(db)}
    rows = []
    snap_ids = set()
    for r in snapshot["rows"]:
        if r["indicator_id"] not in known:
            continue
        snap_ids.add(r["indicator_id"])
        rows.append(RowInput(r["indicator_id"], r["value"], r["observation"] or "",
                             [Decimal(str(d["value"])) for d in r.get("details") or []]))
    # Lignes absentes de l'instantané : remises à « non renseigné ».
    for v in db.execute(select(BilanValue).where(BilanValue.bilan_id == bilan.id)).scalars():
        if v.indicator_id not in snap_ids:
            rows.append(RowInput(v.indicator_id, None, "", []))
    return apply_rows(db, bilan, rows, user_id, source=source, reason=reason)


def clear_draft(db: Session, iso_year: int, iso_week: int, user: User, ip: str | None) -> dict[str, Any]:
    """Brouillon jamais validé : effacé. Bilan rouvert : retour à la dernière version validée."""
    bilan = _locked_existing(db, iso_year, iso_week)
    if bilan.status not in EDITABLE:
        raise BilanError(409, "Ce bilan n'est pas un brouillon.", "not_draft")
    if bilan.current_version > 0:
        version = latest_version(db, bilan.id)
        assert version is not None
        _restore_rows_from_snapshot(db, bilan, version.snapshot, user.id, "restauration",
                                    "Abandon des modifications après réouverture")
        bilan.status = "valide"
        bilan.revision += 1
        audit.record(db, "bilan.reopen_cancelled", user_id=user.id, username=user.username, target_type="bilan",
                     target_id=_week_ref(bilan), ip=ip, details={"version_restauree": version.version_no})
        db.commit()
        return bilan_payload(db, iso_year, iso_week)
    count = len(bilan.values) if bilan.values else 0
    ref = _week_ref(bilan)
    snapshot = snapshot_of(db, bilan)
    db.delete(bilan)
    audit.record(db, "bilan.draft_cleared", user_id=user.id, username=user.username, target_type="bilan",
                 target_id=ref, ip=ip, details={"lignes": count, "contenu_efface": snapshot["rows"]})
    db.commit()
    return bilan_payload(db, iso_year, iso_week)


def copy_previous(db: Session, iso_year: int, iso_week: int, revision: int, with_observations: bool, user: User,
                  ip: str | None) -> dict[str, Any]:
    (py, pw), prev = previous_rows(db, iso_year, iso_week)
    if not prev:
        raise BilanError(404, f"Aucun bilan pour la semaine {pw} de {py}.", "no_previous")
    bilan = get_or_create_for_update(db, iso_year, iso_week, user.id)
    if bilan.status not in EDITABLE:
        raise BilanError(423, "Ce bilan doit être rouvert avant d'être modifié.", "locked")
    if bilan.revision != revision:
        raise BilanError(409, "Le bilan a changé entre-temps. Rechargez la page.", "conflict")
    active = {i.id for i in get_indicators(db, include_inactive=False)}
    rows = [RowInput(r["indicator_id"], r["value"], (r["observation"] or "") if with_observations else "",
                     [Decimal(str(d["value"])) for d in r.get("details") or []] if with_observations else [])
            for r in prev.values() if r["indicator_id"] in active]
    prov = {r.indicator_id: {"locator": f"Semaine {pw} de {py}", "ref": None} for r in rows}
    changed = apply_rows(db, bilan, rows, user.id, source="reprise", provenance=prov,
                         reason=f"Reprise de la semaine {pw}")
    bilan.revision += 1
    bilan.updated_by, bilan.updated_at = user.id, utcnow()
    audit.record(db, "bilan.copied_previous", user_id=user.id, username=user.username, target_type="bilan",
                 target_id=_week_ref(bilan), ip=ip,
                 details={"source": f"{py}-S{pw:02d}", "lignes_modifiees": changed, "observations": with_observations})
    db.commit()
    return bilan_payload(db, iso_year, iso_week)


def restore_version(db: Session, iso_year: int, iso_week: int, version_no: int, user: User,
                    ip: str | None) -> dict[str, Any]:
    bilan = _locked_existing(db, iso_year, iso_week)
    if bilan.status != "rouvert":
        raise BilanError(409, "Rouvrez le bilan avant de restaurer une ancienne version.", "not_reopened")
    version = db.execute(select(BilanVersion).where(BilanVersion.bilan_id == bilan.id,
                                                    BilanVersion.version_no == version_no)).scalar_one_or_none()
    if version is None:
        raise BilanError(404, "Version introuvable.", "not_found")
    changed = _restore_rows_from_snapshot(db, bilan, version.snapshot, user.id, "restauration",
                                          f"Restauration de la version {version_no}")
    bilan.revision += 1
    bilan.updated_by, bilan.updated_at = user.id, utcnow()
    audit.record(db, "bilan.version_restored", user_id=user.id, username=user.username, target_type="bilan",
                 target_id=_week_ref(bilan), ip=ip, details={"version": version_no, "lignes_modifiees": changed})
    db.commit()
    return bilan_payload(db, iso_year, iso_week)


def list_versions(db: Session, iso_year: int, iso_week: int) -> list[dict[str, Any]]:
    bilan = find_bilan(db, iso_year, iso_week)
    if bilan is None:
        return []
    versions = list(db.execute(select(BilanVersion).where(BilanVersion.bilan_id == bilan.id)
                               .order_by(BilanVersion.version_no.desc())).scalars())
    names = _user_names(db, {v.created_by for v in versions})
    return [{"version_no": v.version_no, "created_at": v.created_at.isoformat(),
             "created_by": names.get(v.created_by or 0), "sha256": v.snapshot_sha256, "comment": v.comment,
             "integrity_ok": snapshot_digest(v.snapshot) == v.snapshot_sha256, "snapshot": v.snapshot}
            for v in versions]


def history(db: Session, iso_year: int, iso_week: int, limit: int = 300) -> list[dict[str, Any]]:
    bilan = find_bilan(db, iso_year, iso_week)
    if bilan is None:
        return []
    rows = list(db.execute(select(ValueHistory).where(ValueHistory.bilan_id == bilan.id)
                           .order_by(ValueHistory.id.desc()).limit(limit)).scalars())
    names = _user_names(db, {r.user_id for r in rows})
    labels = {i.id: i.label for i in get_indicators(db)}
    files = _import_names(db, {r.import_id for r in rows})
    return [{"at": r.last_at.isoformat(), "indicator": labels.get(r.indicator_id), "field": r.field,
             "old": r.old_value, "new": r.new_value, "source": r.source_kind,
             "file": files.get(r.import_id) if r.import_id else None, "reason": r.reason,
             "user": names.get(r.user_id or 0)} for r in rows]


def list_bilans(db: Session, iso_year: int | None = None) -> list[dict[str, Any]]:
    stmt = select(Bilan).order_by(Bilan.iso_year.desc(), Bilan.iso_week.desc())
    if iso_year:
        stmt = stmt.where(Bilan.iso_year == iso_year)
    bilans = list(db.execute(stmt.options(selectinload(Bilan.values))).scalars())
    names = _user_names(db, {b.validated_by for b in bilans} | {b.updated_by for b in bilans})
    return [{
        "iso_year": b.iso_year, "iso_week": b.iso_week, "period_start": b.period_start.isoformat(),
        "period_end": b.period_end.isoformat(), "status": b.status, "status_label": STATUS_LABELS[b.status],
        "current_version": b.current_version, "filled": sum(1 for v in b.values if v.value is not None),
        "updated_at": b.updated_at.isoformat(), "updated_by": names.get(b.updated_by or 0),
        "validated_at": b.validated_at.isoformat() if b.validated_at else None,
        "validated_by": names.get(b.validated_by or 0), "is_demo": b.is_demo,
    } for b in bilans]


@dataclass
class DefaultWeek:
    iso_year: int
    iso_week: int
    reason: str
    drafts: list[str] = field(default_factory=list)


def default_week(db: Session) -> DefaultWeek:
    """Semaine affichée à l'arrivée (voir DECISIONS D-21) :
    1. le brouillon en cours le plus récent ;
    2. sinon la dernière semaine écoulée si son bilan n'est pas encore validé (le bilan d'une semaine
       se fait une fois la semaine terminée) ;
    3. sinon la semaine courante."""
    cy, cw = isoweek.current_week()
    drafts = list(db.execute(select(Bilan).where(Bilan.status.in_(EDITABLE))
                             .order_by(Bilan.updated_at.desc())).scalars())
    labels = [f"S{b.iso_week} {b.iso_year}" for b in drafts]
    if drafts:
        b = drafts[0]
        return DefaultWeek(b.iso_year, b.iso_week, "brouillon", labels)
    py, pw = isoweek.shift_week(cy, cw, -1)
    prev = find_bilan(db, py, pw)
    if prev is None or prev.status in EDITABLE:
        return DefaultWeek(py, pw, "semaine_ecoulee", labels)
    return DefaultWeek(cy, cw, "courante", labels)
