#!/usr/bin/env python3
"""Generate reproducible benchmark delta reports for two run summaries."""

from __future__ import annotations

import argparse
import json
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, List


TRUE_VALUES = {
    "1",
    "true",
    "yes",
    "on",
    "passed",
    "pass",
    "ok",
    "success",
}
FALSE_VALUES = {
    "0",
    "false",
    "no",
    "off",
    "failed",
    "fail",
    "error",
    "timed_out",
    "invalid",
    "timeout",
    "unknown",
    "skipped",
}
PASSABLE_STATUSES = {"passed", "pass", "success", "ok"}


def _coerce_runtime(value: Any) -> float:
    try:
        return float(value)
    except (TypeError, ValueError):
        return 0.0


def _coerce_pass(value: Any, *, status: Any = None) -> int:
    if isinstance(value, bool):
        return int(value)

    if isinstance(value, (int, float)):
        try:
            return 1 if float(value) != 0 else 0
        except (TypeError, ValueError):
            return 0

    if value is None:
        if status is not None and str(status).strip().lower() == "passed":
            return 1
        return 0

    if isinstance(value, (bytes, bytearray)):
        try:
            value = value.decode("utf-8")
        except Exception:
            value = ""

    if isinstance(value, str):
        normalized = value.strip().lower()
        if not normalized:
            return 0
        if normalized in TRUE_VALUES:
            return 1
        if normalized in FALSE_VALUES:
            return 0
        try:
            return 1 if float(normalized) != 0 else 0
        except ValueError:
            return 0

    return 1 if bool(value) else 0


def _normalize_status(value: Any) -> str:
    text = "" if value is None else str(value).strip().lower()
    return text or "unknown"


def _status_score(value: Any) -> int:
    return 1 if _normalize_status(value) in PASSABLE_STATUSES else 0


def _normalize_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
    raw = entry if isinstance(entry, dict) else {}
    name = raw.get("name") or raw.get("id") or "unknown"
    status = raw.get("status")
    runtime = _coerce_runtime(raw.get("runtime_sec", 0.0))
    passed = _coerce_pass(raw.get("pass"), status=status)
    normalized = dict(raw)
    normalized["name"] = str(name)
    normalized["status"] = str(status) if status is not None else "unknown"
    normalized["runtime_sec"] = runtime
    normalized["pass"] = passed
    normalized["_status_score"] = _status_score(status)
    if normalized.get("metric") is None:
        normalized["metric"] = {"metric_unavailable": True}
    return normalized


def _load_summary(summary_path: Path) -> Dict[str, Any]:
    payload = json.loads(summary_path.read_text(encoding="utf-8"))
    if not isinstance(payload, dict):
        raise ValueError(f"summary JSON is not an object: {summary_path}")
    raw_benchmarks = payload.get("benchmarks", [])
    if not isinstance(raw_benchmarks, list):
        raw_benchmarks = []
    normalized = [_normalize_entry(b) for b in raw_benchmarks]
    by_name = {str(b["name"]): b for b in normalized}
    return {
        "run_id": str(payload.get("run_id", "unknown")),
        "generated_at": payload.get("generated_at", ""),
        "summary": payload.get("summary", {}),
        "benchmarks": by_name,
        "_benchmark_list": normalized,
    }


def _read_summary(base_runs_dir: Path, run_id: str) -> Dict[str, Any]:
    summary_path = base_runs_dir / run_id / "summary.json"
    if not summary_path.exists():
        raise FileNotFoundError(f"summary.json not found for run_id '{run_id}' at: {summary_path}")
    return _load_summary(summary_path)


def _format_delta(value: float | None, precision: int = 6) -> str:
    if value is None:
        return "n/a"
    return f"{value:+.{precision}f}"


def _float_or_none(value: Any) -> float | None:
    if value is None:
        return None
    try:
        v = float(value)
        if v == int(v) and abs(v) > 0 and v.is_integer():
            return float(f"{v:.6f}")
        return float(f"{v:.6f}")
    except (TypeError, ValueError):
        return None


def generate_delta(baseline_id: str, target_id: str, runs_dir: Path) -> Dict[str, Any]:
    baseline = _read_summary(runs_dir, baseline_id)
    target = _read_summary(runs_dir, target_id)

    baseline_benchmarks: Dict[str, Dict[str, Any]] = baseline["benchmarks"]
    target_benchmarks: Dict[str, Dict[str, Any]] = target["benchmarks"]

    baseline_names = set(baseline_benchmarks)
    target_names = set(target_benchmarks)
    all_names = sorted(baseline_names | target_names)

    compared_names = sorted(baseline_names & target_names)

    benchmark_rows: List[Dict[str, Any]] = []
    status_improved = 0
    status_regressed = 0
    runtime_improved = 0
    runtime_regressed = 0
    runtime_missing_count = 0
    runtime_delta_sum = 0.0
    runtime_delta_count = 0
    pass_delta_sum = 0

    for name in all_names:
        base = baseline_benchmarks.get(name)
        run = target_benchmarks.get(name)

        base_status = base["status"] if base else None
        run_status = run["status"] if run else None
        base_pass = base["pass"] if base else 0
        run_pass = run["pass"] if run else 0
        base_runtime = base["runtime_sec"] if base else None
        run_runtime = run["runtime_sec"] if run else None

        status_delta = (run["_status_score"] if run else 0) - (base["_status_score"] if base else 0)
        pass_delta = run_pass - base_pass

        if base is not None and run is not None:
            runtime_delta = run_runtime - base_runtime
            runtime_delta_percent = (runtime_delta / base_runtime * 100.0) if base_runtime else None
            pass_delta_sum += pass_delta
            runtime_delta_sum += runtime_delta
            runtime_delta_count += 1
            if status_delta > 0:
                status_improved += 1
            elif status_delta < 0:
                status_regressed += 1

            if runtime_delta < 0:
                runtime_improved += 1
            elif runtime_delta > 0:
                runtime_regressed += 1
        else:
            runtime_missing_count += 1
            runtime_delta = run_runtime - base_runtime if (base_runtime is not None and run_runtime is not None) else None
            if base_runtime is None or run_runtime is None:
                runtime_delta_percent = None

        benchmark_rows.append(
            {
                "name": name,
                "baseline": {
                    "status": base_status,
                    "pass": base_pass,
                    "runtime_sec": base_runtime,
                    "metric": base.get("metric") if base else None,
                },
                "target": {
                    "status": run_status,
                    "pass": run_pass,
                    "runtime_sec": run_runtime,
                    "metric": run.get("metric") if run else None,
                },
                "status_change": f"{_normalize_status(base_status)} -> {_normalize_status(run_status)}",
                "status_delta": int(status_delta),
                "pass_delta": int(pass_delta),
                "runtime_delta": _float_or_none(runtime_delta),
                "runtime_delta_percent": _float_or_none(runtime_delta_percent),
                "included_in_common": bool(base is not None and run is not None),
            }
        )

    runtime_delta_avg = runtime_delta_sum / runtime_delta_count if runtime_delta_count else 0.0

    baseline_pass_total = sum(int(v["pass"]) for v in baseline_benchmarks.values())
    target_pass_total = sum(int(v["pass"]) for v in target_benchmarks.values())

    return {
        "generated_at": datetime.now(tz=timezone.utc).isoformat(),
        "baseline_run_id": baseline_id,
        "target_run_id": target_id,
        "totals": {
            "baseline_benchmarks": len(baseline_names),
            "target_benchmarks": len(target_names),
            "compared_benchmarks": len(compared_names),
            "added_benchmarks": len(target_names - baseline_names),
            "removed_benchmarks": len(baseline_names - target_names),
            "runtime_delta_count": runtime_delta_count,
            "runtime_delta_sum_sec": round(runtime_delta_sum, 6),
            "runtime_delta_avg_sec": round(runtime_delta_avg, 6),
            "runtime_improvements": runtime_improved,
            "runtime_regressions": runtime_regressed,
            "runtime_missing": runtime_missing_count,
            "status_improvements": status_improved,
            "status_regressions": status_regressed,
            "baseline_pass_total": int(baseline_pass_total),
            "target_pass_total": int(target_pass_total),
            "pass_delta_total": int(target_pass_total - baseline_pass_total),
            "common_pass_delta_sum": int(pass_delta_sum),
        },
        "benchmarks": benchmark_rows,
    }


def _write_markdown(delta: Dict[str, Any], path: Path) -> None:
    totals = delta["totals"]
    lines = [
        "# Benchmark Delta Report",
        "",
        f"- Baseline run: `{delta['baseline_run_id']}`",
        f"- Target run: `{delta['target_run_id']}`",
        f"- Generated at: `{delta['generated_at']}`",
        "",
        "## Aggregate counts",
        "",
        f"- Baseline benchmarks: {totals['baseline_benchmarks']}",
        f"- Target benchmarks: {totals['target_benchmarks']}",
        f"- Compared benchmarks: {totals['compared_benchmarks']}",
        f"- Added benchmarks: {totals['added_benchmarks']}",
        f"- Removed benchmarks: {totals['removed_benchmarks']}",
        f"- Runtime delta (sum): {totals['runtime_delta_sum_sec']} sec",
        f"- Runtime delta (avg over compared): {totals['runtime_delta_avg_sec']} sec",
        f"- Runtime improvements: {totals['runtime_improvements']}",
        f"- Runtime regressions: {totals['runtime_regressions']}",
        f"- Status improvements: {totals['status_improvements']}",
        f"- Status regressions: {totals['status_regressions']}",
        f"- Baseline pass total: {totals['baseline_pass_total']}",
        f"- Target pass total: {totals['target_pass_total']}",
        f"- Net pass delta: {totals['pass_delta_total']}",
        f"",
        "## Per-benchmark deltas",
        "",
        "|Benchmark|Baseline status|Target status|Status delta|Pass (base -> target)|Pass delta|Runtime delta (sec)|Runtime delta (%)|",
        "|---|---|---|---:|---:|---:|---:|---:|",
    ]

    for row in delta["benchmarks"]:
        base = row["baseline"]
        tar = row["target"]
        base_status = base["status"] or "missing"
        target_status = tar["status"] or "missing"
        runtime_delta = row["runtime_delta"]
        runtime_pct = row["runtime_delta_percent"]
        lines.append(
            "|{name}|{base_status}|{target_status}|{status_delta}|{base_pass} -> {target_pass}|{pass_delta}|{runtime_delta}|{runtime_pct}|".format(
                name=row["name"],
                base_status=base_status,
                target_status=target_status,
                status_delta=row["status_delta"],
                base_pass=base["pass"],
                target_pass=tar["pass"],
                pass_delta=_format_delta(row["pass_delta"], precision=0),
                runtime_delta=_format_delta(runtime_delta),
                runtime_pct=(f"{runtime_pct:+.2f}%" if runtime_pct is not None else "n/a"),
            )
        )

    path.write_text("\n".join(lines) + "\n", encoding="utf-8")


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Generate benchmark delta reports between two run directories.")
    parser.add_argument("--baseline-run", required=True, help="Baseline run id, e.g. 2026-03-04-baseline-refresh")
    parser.add_argument("--target-run", required=True, help="Target run id, e.g. 2026-03-07-rebenchmark-delta")
    parser.add_argument("--runs-dir", default="ai_research/benchmarks/runs", help="Directory containing run folders")
    parser.add_argument("--output-dir", default="ai_research/benchmarks", help="Directory for generated DELTA_*.md/json")
    return parser.parse_args()


def main() -> int:
    args = parse_args()

    runs_dir = Path(args.runs_dir)
    out_dir = Path(args.output_dir)
    out_dir.mkdir(parents=True, exist_ok=True)

    delta = generate_delta(args.baseline_run, args.target_run, runs_dir)
    output_stem = f"DELTA_{args.target_run}"

    json_path = out_dir / f"{output_stem}.json"
    md_path = out_dir / f"{output_stem}.md"

    json_path.write_text(json.dumps(delta, indent=2), encoding="utf-8")
    _write_markdown(delta, md_path)

    print(json_path)
    print(md_path)
    return 0


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