#!/usr/bin/env python3
"""KAR-755 — Distributional Trace Analytics über aria-llm-trace.py-Traces.

Liest /root/aria/logs/llm-traces.jsonl (KAR-744) und erzeugt einen Drift-/
Distributions-Report mit ECHTEN Zahlen: pro (script, model) Latenz-/Token-/
Error-Verteilungen, Triage-Score-Distributionen + Verdict-Skew, Model-Pin-Drift.

Punkt-Evals (aria-eval-judge) messen einzelne Outputs; dieser Layer misst die
VERTEILUNG über viele Runs — findet stille Drift + Failure-Modes die Evals missen.

Pure stdlib. Wirft nicht bei kaputten Zeilen.

Usage:
  aria-trace-analytics.py [--input PATH] [--out REPORT.md] [--ref-model MODEL]
"""
from __future__ import annotations
import argparse
import json
import re
import statistics as st
from collections import Counter, defaultdict
from datetime import datetime, timezone
from pathlib import Path

TRACE_PATH = Path("/root/aria/logs/llm-traces.jsonl")
TRIAGE_AXES = ["depth", "novelty", "aria_relevance", "production_ready", "junk_penalty"]

# Arias aktuell beabsichtigtes Default-Modell (Memory reference_fable_mythos_suspension_2026_06)
CURRENT_DEFAULT_MODEL = "claude-opus-4-8"


def _load(path: Path) -> list[dict]:
    rows = []
    if not path.exists():
        return rows
    for line in path.open(encoding="utf-8"):
        line = line.strip()
        if not line:
            continue
        try:
            rows.append(json.loads(line))
        except json.JSONDecodeError:
            continue
    return rows


def _pct(values: list[float], p: float):
    if not values:
        return None
    s = sorted(values)
    k = max(0, min(len(s) - 1, int(round((p / 100) * (len(s) - 1)))))
    return s[k]


def _parse_json_blob(text):
    """Extrahiert das erste JSON-Objekt aus einer (ggf. ```json-gefencten) Response."""
    if not text:
        return None
    m = re.search(r"```(?:json)?\s*(\{.*?\})\s*```", text, re.S)
    blob = m.group(1) if m else None
    if blob is None:
        m2 = re.search(r"(\{.*\})", text, re.S)
        blob = m2.group(1) if m2 else None
    if not blob:
        return None
    try:
        return json.loads(blob)
    except json.JSONDecodeError:
        return None


def _num(x):
    try:
        return float(x)
    except (TypeError, ValueError):
        return None


def analyse(rows: list[dict]) -> dict:
    by_key = defaultdict(list)  # (script, model) -> rows
    for r in rows:
        by_key[(r.get("script"), r.get("model"))].append(r)

    out = {"groups": [], "triage": None, "total": len(rows), "model_pins": {}}
    for (script, model), rs in sorted(by_key.items(), key=lambda kv: (-len(kv[1]),)):
        lat = [v for v in (_num(r.get("latency_ms")) for r in rs) if v is not None]
        it = [v for v in (_num(r.get("input_tokens")) for r in rs) if v is not None]
        ot = [v for v in (_num(r.get("output_tokens")) for r in rs) if v is not None]
        errs = sum(1 for r in rs if r.get("error"))
        ts = sorted(r.get("ts", "") for r in rs if r.get("ts"))
        out["groups"].append({
            "script": script, "model": model, "n": len(rs),
            "ts_from": ts[0][:19] if ts else "?", "ts_to": ts[-1][:19] if ts else "?",
            "lat_p50": _pct(lat, 50), "lat_p95": _pct(lat, 95), "lat_max": max(lat) if lat else None,
            "in_med": st.median(it) if it else None, "out_med": st.median(ot) if ot else None,
            "err": errs, "err_rate": round(errs / len(rs), 4),
        })
        out["model_pins"].setdefault(script, set()).add(model)

    # Triage-Score-Distributionen
    triage = [r for r in rows if r.get("script") == "aria-akp-triage" and r.get("response")]
    if triage:
        axes = {a: Counter() for a in TRIAGE_AXES}
        verdicts = Counter()
        parsed = 0
        for r in triage:
            d = _parse_json_blob(r.get("response"))
            if not d:
                continue
            parsed += 1
            for a in TRIAGE_AXES:
                if a in d:
                    axes[a][d[a]] += 1
            if "verdict" in d:
                verdicts[d["verdict"]] += 1
        out["triage"] = {
            "n": len(triage), "parsed": parsed,
            "axes": {a: dict(sorted(c.items())) for a, c in axes.items()},
            "verdicts": dict(verdicts.most_common()),
            "axis_means": {a: round(sum(int(k) * v for k, v in c.items()) / sum(c.values()), 2)
                           for a, c in axes.items() if c},
        }
    return out


def render(a: dict, ref_model: str) -> str:
    now = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M UTC")
    L = []
    L.append(f"# LLM-Trace Distributional Analytics (KAR-755)\n")
    L.append(f"> Generiert {now} aus `{TRACE_PATH}` — {a['total']} Traces. Echte Zahlen, kein Schätzwert.\n")

    L.append("## Pro Script/Model\n")
    L.append("| Script | Model | n | Zeitraum | Lat p50 | Lat p95 | Lat max | In(med) | Out(med) | Err-Rate |")
    L.append("|---|---|---|---|---|---|---|---|---|---|")
    for g in a["groups"]:
        L.append(f"| {g['script']} | {g['model']} | {g['n']} | {g['ts_from'][:10]}→{g['ts_to'][:10]} | "
                 f"{g['lat_p50']} | {g['lat_p95']} | {g['lat_max']} | {g['in_med']} | {g['out_med']} | {g['err_rate']} |")
    L.append("")

    # Model-Pin-Drift
    L.append("## Model-Pin-Drift\n")
    flagged = []
    for script, models in a["model_pins"].items():
        for m in models:
            if m and ref_model not in m and "haiku" not in (m or "") and "sonnet" not in (m or ""):
                # Opus-Familie aber nicht das aktuelle Default
                if "opus" in m and m != ref_model:
                    flagged.append((script, m))
    if flagged:
        for script, m in flagged:
            L.append(f"- ⚠ `{script}` nutzt `{m}` — aktuelles Default ist `{ref_model}`. Pin prüfen (stille Drift nach Modell-Wechsel).")
    else:
        L.append("- Keine Opus-Pin-Abweichung vom aktuellen Default gefunden.")
    L.append("")

    # Triage-Verteilungen
    if a["triage"]:
        t = a["triage"]
        L.append("## Triage-Judge Score-Distributionen\n")
        L.append(f"Basis: {t['n']} Triage-Traces, {t['parsed']} JSON-parsebar.\n")
        L.append("**Verdict-Verteilung:**\n")
        tot = sum(t["verdicts"].values()) or 1
        for v, c in t["verdicts"].items():
            L.append(f"- `{v}`: {c} ({round(100*c/tot)}%)")
        L.append("")
        L.append("**Score-Achsen (Verteilung 0-3 + Mittel):**\n")
        L.append("| Achse | 0 | 1 | 2 | 3 | Mittel |")
        L.append("|---|---|---|---|---|---|")
        for ax in TRIAGE_AXES:
            dist = t["axes"].get(ax, {})
            row = " | ".join(str(dist.get(i, 0)) for i in range(4))
            L.append(f"| {ax} | {row} | {t['axis_means'].get(ax,'—')} |")
        L.append("")
        L.append("> Diese Verteilung ist die **Baseline** für n-Run-Distribution-Diffs nach Prompt-/Modell-Änderung.")
        L.append("")

    L.append("## Nicht abgedeckt\n")
    L.append("- **Lazy-Tool-Call-Detector** (declared vs executed): braucht Claude-Code-Tool-Call-Traces "
             "(`agent-trace-export`), nicht LLM-Content-Traces. Offen bis agent-trace-export aktiv (KAR-744 Scope-3).")
    L.append("- **Topic-Clustering** (Embeddings): optional, separater Lauf.")
    return "\n".join(L) + "\n"


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--input", default=str(TRACE_PATH))
    ap.add_argument("--out", default=None)
    ap.add_argument("--ref-model", default=CURRENT_DEFAULT_MODEL)
    args = ap.parse_args()
    rows = _load(Path(args.input))
    if not rows:
        print("Keine Traces gefunden.")
        return 1
    report = render(analyse(rows), args.ref_model)
    if args.out:
        Path(args.out).write_text(report, encoding="utf-8")
        print(f"Report geschrieben: {args.out} ({len(rows)} Traces)")
    else:
        print(report)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
