#!/usr/bin/env python3
# ruff: noqa: E501
"""Compare two Codex benchmark reports using deterministic distributions."""
from __future__ import annotations

import argparse
import json
import statistics
import sys
from pathlib import Path
from typing import Any


METRICS = {
    "tokens.total": ("tokens", "total"),
    "tokens.input": ("tokens", "input"),
    "tokens.uncached_input": ("tokens", "uncached_input"),
    "tokens.output": ("tokens", "output"),
    "tokens.reasoning_output": ("tokens", "reasoning_output"),
    "root_user_turns": ("metrics", "user_turns"),
    "discovery_calls": ("metrics", "discovery_calls"),
    "repeated_discovery_count": ("derived", "repeated_discovery_count"),
    "compactions": ("metrics", "compactions"),
    "qa_invocations": ("metrics", "qa_invocations"),
    "full_qa_rate": ("derived", "full_qa_rate"),
    "qa_repair_cycles": ("derived", "qa_repair_cycles"),
    "post_review_rework_cycles": ("derived", "post_review_rework_cycles"),
    "subagent_token_share": ("derived", "subagent_token_share"),
    "duration_seconds": (None, "duration_seconds"),
}


def percentile(values: list[float], fraction: float) -> float | None:
    """Return linear-interpolated percentile (the NumPy/R7 convention)."""
    if not values:
        return None
    ordered = sorted(values)
    if len(ordered) == 1:
        return ordered[0]
    position = (len(ordered) - 1) * fraction
    lower = int(position)
    upper = min(lower + 1, len(ordered) - 1)
    return ordered[lower] + (ordered[upper] - ordered[lower]) * (position - lower)


def _values(report: dict[str, Any], location: tuple[str | None, str]) -> list[float]:
    parent, key = location
    result: list[float] = []
    for session in report.get("sessions", []):
        if not isinstance(session, dict):
            continue
        value: Any = session.get(key) if parent is None else session.get(parent, {}).get(key)
        if isinstance(value, bool) or not isinstance(value, (int, float)):
            continue
        result.append(float(value))
    return result


def _distribution(before: list[float], after: list[float]) -> dict[str, Any]:
    paired = min(len(before), len(after))
    before = before[:paired]
    after = after[:paired]
    before_median = statistics.median(before) if before else None
    after_median = statistics.median(after) if after else None
    delta = None if before_median in (None, 0) or after_median is None else (after_median - before_median) / before_median * 100
    return {
        "count": paired,
        "before": {"median": before_median, "p75": percentile(before, 0.75)},
        "after": {"median": after_median, "p75": percentile(after, 0.75)},
        "percent_delta": delta,
    }


def compare(baseline: dict[str, Any], after: dict[str, Any]) -> dict[str, Any]:
    return {
        "schema_version": "1.0",
        "codex": {name: _distribution(_values(baseline, location), _values(after, location)) for name, location in METRICS.items()},
        "docker_agent_overhead": {"token_usage": None, "budget_ceiling_tokens": 20000, "measurement_status": "unavailable"},
    }


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("baseline", type=Path)
    parser.add_argument("after", type=Path)
    parser.add_argument("--output", type=Path, required=True)
    args = parser.parse_args()
    try:
        baseline = json.loads(args.baseline.read_text(encoding="utf-8"))
        after = json.loads(args.after.read_text(encoding="utf-8"))
        if not isinstance(baseline, dict) or not isinstance(after, dict):
            raise ValueError("reports must contain JSON objects")
        result = compare(baseline, after)
        args.output.parent.mkdir(parents=True, exist_ok=True)
        args.output.write_text(json.dumps(result, indent=2, sort_keys=True) + "\n", encoding="utf-8")
    except (OSError, ValueError, json.JSONDecodeError) as exc:
        parser.error(str(exc))
    return 0


if __name__ == "__main__":
    sys.exit(main())
