#!/usr/bin/env python3
"""Generate benchmark class trend artifacts from historical benchmark run summaries.

This script reads run summaries from:
  ai_research/benchmarks/runs/*/summary.json
and emits:
  ai_research/benchmarks/CLASS_TRENDS.json
  ai_research/benchmarks/CLASS_TRENDS.md

It is intentionally deterministic:
- run selection is stable-sorted by date then run id (descending)
- class keys are sorted deterministically
- JSON output is dumped with sort_keys=True
- all computed derived fields are order-independent except for explicitly-specified ordering
"""

from __future__ import annotations

import argparse
import hashlib
import json
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List, Tuple


ROOT = Path(__file__).resolve().parent
DEFAULT_RUNS_DIR = ROOT / "runs"
DEFAULT_OUTPUT_JSON = ROOT / "CLASS_TRENDS.json"
DEFAULT_OUTPUT_MD = ROOT / "CLASS_TRENDS.md"


@dataclass(frozen=True)
class RunRecord:
    run_id: str
    date_key: str  # YYYY-MM-DD string used for sorting
    date_sort: str  # ISO-ish sortable timestamp string
    source: str
    classes: Dict[str, Dict[str, Any]]


def _read_json(path: Path) -> Any:
    with path.open("r", encoding="utf-8") as f:
        return json.load(f)


def _extract_run_id(payload: Dict[str, Any], fallback: Path) -> str:
    for key in ("run_id", "name", "run", "id", "run_name"):
        val = payload.get(key)
        if isinstance(val, str) and val.strip():
            return val
    return fallback.parent.name


def _coerce_date_text(value: Any) -> str:
    """Return a sortable timestamp-like string and a display date."""
    if not isinstance(value, str):
        return ""
    v = value.strip()
    if not v:
        return ""
    # Accept already-compact date-only, ISO8601, or plain date strings.
    return v


def _parse_run_date_for_sort(value: str) -> str:
    """Normalize to a stable sortable key.

    If parsing fails, return the original text.
    """
    if not value:
        return ""
    # Most generated_at values are ISO timestamps, e.g. 2026-03-11T07:19:28Z
    for fmts in [
        "%Y-%m-%dT%H:%M:%SZ",
        "%Y-%m-%dT%H:%M:%S",
        "%Y-%m-%d",
    ]:
        try:
            dt = datetime.strptime(value[: len(fmts)], fmts)
            return dt.isoformat()
        except Exception:
            pass
    return value


def _extract_run_date(payload: Dict[str, Any], summary_path: Path) -> Tuple[str, str]:
    """(date_key, date_sort_key)."""
    # Preferred keys from known/legacy contracts.
    for key in (
        "generated_at",
        "timestamp",
        "timestamp_utc",
        "started_at",
        "created_at",
        "date",
        "run_date",
    ):
        v = _coerce_date_text(payload.get(key))
        if v:
            # If generated_at is full timestamp, use date prefix as display date.
            return (v[:10], _parse_run_date_for_sort(v))

    # Fallback to folder names that often encode run date.
    fallback = summary_path.parent.name
    return (fallback[:10], _parse_run_date_for_sort(fallback))


def _to_float(value: Any, default: float = 0.0) -> float:
    try:
        return float(value)
    except Exception:
        return default


def _to_int(value: Any, default: int = 0) -> int:
    try:
        return int(value)
    except Exception:
        return default


def _normalize_benchmark_to_class_entry(bench: Dict[str, Any]) -> Tuple[str, Dict[str, Any]] | None:
    """Return (class_name, class_metrics)."""
    if not isinstance(bench, dict):
        return None

    class_name = None
    # Deterministic class identity: prefer stable id, then name.
    cid = bench.get("id")
    if isinstance(cid, str) and cid.strip():
        class_name = cid.strip()
    if not class_name:
        n = bench.get("name")
        if isinstance(n, str) and n.strip():
            class_name = n.strip()

    if not class_name:
        return None

    passed = _to_int(bench.get("pass"), default=0)
    status = bench.get("status")
    status_text = str(status).strip().lower() if status is not None else "unknown"
    runtime = _to_float(bench.get("runtime_sec"), default=0.0)

    # Treat pass indicator as class score and pass-rate source.
    pass_rate = _to_float(bench.get("pass_rate"), default=_to_float(passed, default=0.0))
    # If pass_rate looks like fraction-percentage (e.g., 100), normalize to [0,1]
    if pass_rate > 1.0:
        pass_rate = pass_rate / 100.0

    return class_name, {
        "run_count": 1,
        "passed": passed,
        "score": _to_float(passed, default=0.0),
        "pass_rate": pass_rate,
        "status": status_text,
        "runtime_sec": runtime,
        "bench_id": str(bench.get("id") or ""),
        "bench_name": str(bench.get("name") or ""),
    }


def _extract_class_results(summary: Dict[str, Any]) -> Dict[str, Dict[str, Any]]:
    """Map class name to aggregate metrics for the run.

    Current repository summaries use:
      - summary['benchmarks']: list of benchmark results

    This normalizer tolerates absent or mixed keys for forward compatibility.
    """
    benchmarks = summary.get("benchmarks")
    if isinstance(benchmarks, dict):
        # Legacy/alternate shape support.
        entries = []
        for k, v in benchmarks.items():
            if isinstance(v, dict):
                item = dict(v)
                if not item.get("id"):
                    item["id"] = k
                entries.append(item)
            else:
                entries.append({"id": k, "pass": _to_int(v), "status": "passed" if _to_int(v) else "failed", "pass_rate": _to_float(v)})
    elif isinstance(benchmarks, list):
        entries = benchmarks
    else:
        entries = []

    class_results: Dict[str, Dict[str, Any]] = {}
    for bench in entries:
        norm = _normalize_benchmark_to_class_entry(bench if isinstance(bench, dict) else {})
        if not norm:
            continue
        cname, metrics = norm
        # In current shape, class appears once per run. Keep deterministic single entry.
        class_results[cname] = metrics

    return class_results


def _load_runs(runs_root: Path) -> List[RunRecord]:
    runs: List[RunRecord] = []

    for path in sorted(runs_root.glob("*/summary.json")):
        try:
            payload = _read_json(path)
            if not isinstance(payload, dict):
                continue
            run_id = _extract_run_id(payload, path)
            date_key, date_sort = _extract_run_date(payload, path)
            classes = _extract_class_results(payload)
            if not classes:
                continue
            runs.append(
                RunRecord(
                    run_id=run_id,
                    date_key=date_key,
                    date_sort=date_sort,
                    source=str(path),
                    classes=classes,
                )
            )
        except Exception:
            continue

    def _sort_key(r: RunRecord):
        # deterministic latest-first; secondary tie-breaker by run_id
        return (r.date_sort or "", r.run_id)

    runs.sort(key=_sort_key, reverse=True)
    return runs


def _compute_trends(runs: List[RunRecord], latest: int = 3) -> Dict[str, Any]:
    window = max(1, int(latest))
    latest_runs = runs[:window]

    class_timeline: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
    for run in latest_runs:
        for cname, metrics in run.classes.items():
            class_timeline[cname].append(
                {
                    "run_id": run.run_id,
                    "date": run.date_key,
                    "score": metrics.get("score", 0.0),
                    "pass_rate": metrics.get("pass_rate", 0.0),
                    "status": metrics.get("status", "unknown"),
                    "runtime_sec": metrics.get("runtime_sec", 0.0),
                }
            )

    class_trends: Dict[str, Any] = {}
    for cname in sorted(class_timeline):
        points = class_timeline[cname]
        if not points:
            continue
        latest_point = points[0]
        baseline_point = points[-1]
        delta = (latest_point["score"] - baseline_point["score"]) if baseline_point else 0.0
        pass_rate_trend = [
            {
                "run_id": p["run_id"],
                "date": p["date"],
                "pass_rate": p["pass_rate"],
            }
            for p in points
        ]
        class_trends[cname] = {
            "latest": latest_point,
            "baseline": baseline_point,
            "delta": delta,
            "pass_rate_trend": pass_rate_trend,
        }

    return {
        "generated_at": datetime.utcnow().isoformat() + "Z",
        "runs_considered": [
            {
                "index": i + 1,
                "run_id": r.run_id,
                "date": r.date_key,
                "source": r.source,
            }
            for i, r in enumerate(latest_runs)
        ],
        "totals": {
            "runs_available": len(runs),
            "window_requested": window,
            "runs_used": len(latest_runs),
            "classes_seen": len(class_timeline),
        },
        "class_trends": class_trends,
    }


def _to_md_table_row(label: str, latest: Dict[str, Any], baseline: Dict[str, Any], delta: float) -> str:
    latest_score = _to_float(latest.get("score", 0.0))
    baseline_score = _to_float(baseline.get("score", 0.0))
    latest_pass = _to_float(latest.get("pass_rate", 0.0))
    baseline_pass = _to_float(baseline.get("pass_rate", 0.0))
    return (
        f"| {label} "
        f"| {latest_score:.4f} "
        f"| {baseline_score:.4f} "
        f"| {delta:.4f} "
        f"| {latest_pass:.4f} "
        f"| {baseline_pass:.4f} |"
    )


def _build_markdown(payload: Dict[str, Any]) -> str:
    lines: List[str] = []
    totals = payload.get("totals", {})
    runs_considered = payload.get("runs_considered", [])
    class_trends: Dict[str, Any] = payload.get("class_trends", {})

    lines.append("# Benchmark Class Trends")
    lines.append("")
    lines.append("Generated from: `runs/*/summary.json`")
    lines.append(f"Runs available: {totals.get('runs_available', 0)}; window: {totals.get('window_requested', 0)}; runs used: {totals.get('runs_used', 0)}")
    lines.append("")

    if runs_considered:
        lines.append("## Runs considered (newest → oldest)")
        lines.append("| idx | run_id | date | source |");
        lines.append("|---:|---|---|---|")
        for r in runs_considered:
            lines.append(f"| {r['index']} | {r['run_id']} | {r['date']} | {r['source']} |")
        lines.append("")

    lines.append("## Class trend snapshot")
    lines.append("| class | latest score | baseline score | delta | latest pass-rate | baseline pass-rate |")
    lines.append("|---|---:|---:|---:|---:|---:|")

    for cname in sorted(class_trends):
        trend = class_trends[cname]
        lines.append(_to_md_table_row(cname, trend["latest"], trend["baseline"], trend["delta"]))

    lines.append("")
    lines.append("## Pass-rate trends")
    for cname in sorted(class_trends):
        trend = class_trends[cname]["pass_rate_trend"]
        lines.append(f"### {cname}")
        lines.append("| run_id | date | pass_rate |")
        lines.append("|---|---|---:|")
        for point in trend:
            lines.append(f"| {point['run_id']} | {point['date']} | {point['pass_rate']:.4f} |")

        vals = [p["pass_rate"] for p in trend]
        if vals:
            latest_v = vals[0]
            oldest_v = vals[-1]
            avg = sum(vals) / len(vals)
            lines.append(f"\nPass-rate delta (latest - oldest): `{latest_v - oldest_v:.4f}`")
            lines.append(f"Pass-rate average over window: `{avg:.4f}`")
        lines.append("")

    digest = hashlib.sha256(
        json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8")
    ).hexdigest()
    lines.append(f"Deterministic report hash: `{digest}`")
    lines.append("")

    return "\n".join(lines).rstrip() + "\n"


def _write_outputs(payload: Dict[str, Any], out_json: Path, out_md: Path) -> None:
    out_json.parent.mkdir(parents=True, exist_ok=True)
    out_md.parent.mkdir(parents=True, exist_ok=True)
    out_json.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8")
    out_md.write_text(_build_markdown(payload), encoding="utf-8")


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(description="Generate benchmark class trend artifacts")
    parser.add_argument("--runs-dir", default=str(DEFAULT_RUNS_DIR), help="runs directory containing */summary.json")
    parser.add_argument("--latest", type=int, default=3, help="number of latest runs to include")
    parser.add_argument("--out-json", default=str(DEFAULT_OUTPUT_JSON), help="output json path")
    parser.add_argument("--out-md", default=str(DEFAULT_OUTPUT_MD), help="output markdown path")
    args = parser.parse_args(argv)

    runs = _load_runs(Path(args.runs_dir))
    payload = _compute_trends(runs, latest=args.latest)
    _write_outputs(payload, Path(args.out_json), Path(args.out_md))
    return 0


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