"""Journal d'audit en ajout seul, chaîné par SHA-256 (toute altération devient détectable)."""
from __future__ import annotations

import hashlib
import json
from datetime import datetime, timezone
from typing import Any

from sqlalchemy import select, text
from sqlalchemy.orm import Session

from ..models import AuditLog

GENESIS = "0" * 64
_FORBIDDEN_KEYS = {"password", "new_password", "current_password", "secret", "token", "cookie", "code",
                   "mfa_secret", "recovery_code", "csrf"}


def _clean(details: dict[str, Any]) -> dict[str, Any]:
    out: dict[str, Any] = {}
    for k, v in details.items():
        if k.lower() in _FORBIDDEN_KEYS:
            continue
        if isinstance(v, dict):
            v = _clean(v)
        elif isinstance(v, str) and len(v) > 2000:
            v = v[:2000] + "…"
        out[k] = v
    return out


def _digest(prev_hash: str, at: datetime, user_id: int | None, action: str, target_type: str | None,
            target_id: str | None, details: dict[str, Any]) -> str:
    payload = json.dumps(
        [prev_hash, at.astimezone(timezone.utc).isoformat(), user_id, action, target_type, target_id, details],
        sort_keys=True, ensure_ascii=False, separators=(",", ":"), default=str,
    )
    return hashlib.sha256(payload.encode()).hexdigest()


def record(
    db: Session,
    action: str,
    *,
    user_id: int | None = None,
    username: str | None = None,
    target_type: str | None = None,
    target_id: Any = None,
    ip: str | None = None,
    details: dict[str, Any] | None = None,
) -> None:
    """Ajoute une entrée dans la transaction courante (l'appelant valide ou annule)."""
    # Verrou transactionnel : sérialise le chaînage sans bloquer les lectures.
    db.execute(text("SELECT pg_advisory_xact_lock(520052)"))
    prev = db.execute(select(AuditLog.hash).order_by(AuditLog.id.desc()).limit(1)).scalar()
    at = datetime.now(timezone.utc).replace(microsecond=0)
    clean = _clean(details or {})
    tid = None if target_id is None else str(target_id)
    entry = AuditLog(
        at=at, user_id=user_id, username=username, action=action, target_type=target_type,
        target_id=tid, ip=ip, details=clean, prev_hash=prev or GENESIS,
        hash=_digest(prev or GENESIS, at, user_id, action, target_type, tid, clean),
    )
    db.add(entry)
    db.flush()


def verify_chain(db: Session) -> tuple[bool, int, int | None]:
    """Recalcule la chaîne. Renvoie (intègre, nombre de lignes, id de la première ligne fautive)."""
    prev = GENESIS
    count = 0
    for row in db.execute(select(AuditLog).order_by(AuditLog.id)).scalars():
        count += 1
        expected = _digest(prev, row.at, row.user_id, row.action, row.target_type, row.target_id, row.details)
        if row.prev_hash != prev or row.hash != expected:
            return False, count, row.id
        prev = row.hash
    return True, count, None
