#!/usr/bin/env python3
"""
aria-eval-judge — CLI fuer eval-driven-agent-dev Skill (KAR-75)

Laedt golden-dataset.yaml + judge-prompt.md, faehrt jeden Case durch den Judge
(LLM via Anthropic SDK), berechnet Precision/Recall, schreibt Run-JSON nach
<dataset-dir>/runs/<date>.json.

Usage:
  python3 aria-eval-judge.py \
    --dataset /root/aria/evals/<feature>/golden-dataset.yaml \
    --judge-prompt /root/aria/evals/<feature>/judge-prompt.md \
    [--model claude-haiku-4-5-20251001] \
    [--max-cases 10]
"""
from __future__ import annotations
import argparse, json, os, sys, time, yaml
from datetime import datetime, timezone
from pathlib import Path

_ARIA_CONFIG_PATH = Path("/root/aria/brain/youtube/.config.yaml")


def _get_default_model() -> str:
    """Liest eval-Judge-Modell aus ARIA_EVAL_MODEL env, dann .config.yaml models.triage_model,
    dann fällt zurück auf claude-haiku-4-5-20251001."""
    env_val = os.environ.get("ARIA_EVAL_MODEL")
    if env_val:
        return env_val
    try:
        cfg = yaml.safe_load(_ARIA_CONFIG_PATH.read_text(encoding="utf-8"))
        model = cfg.get("models", {}).get("triage_model")
        if model:
            return model
    except Exception:
        pass
    return "claude-haiku-4-5-20251001"


def load_yaml(p: Path) -> dict:
    with p.open() as f:
        return yaml.safe_load(f)


def render_prompt(template: str, case: dict) -> str:
    """Replace {{input}}, {{expected}} placeholders with case data as JSON strings."""
    input_json = json.dumps(case.get("input", {}), indent=2, ensure_ascii=False)
    expected_json = json.dumps(case.get("expected", {}), indent=2, ensure_ascii=False)
    return (
        template
        .replace("{{input_text_or_json}}", input_json)
        .replace("{{input}}", input_json)
        .replace("{{expected_verdict_and_reasons}}", expected_json)
        .replace("{{expected}}", expected_json)
        .replace("{{actual_output}}", "(see input — actual = input for unit-eval mode)")
        .replace("{{actual}}", "(see input — actual = input for unit-eval mode)")
    )


_JUDGE_SYSTEM_PROMPT = (
    "Du bist ein strenger, objektiver Eval-Judge fuer KI-Agenten-Outputs. "
    "Antworte ausschliesslich als JSON mit den Feldern: verdict, scores, reason, confidence."
)


def call_judge(client, model: str, prompt: str) -> tuple[dict | None, dict]:
    try:
        resp = client.messages.create(
            model=model,
            max_tokens=300,
            system=[
                {
                    "type": "text",
                    "text": _JUDGE_SYSTEM_PROMPT,
                    "cache_control": {"type": "ephemeral"},
                }
            ],
            messages=[{"role": "user", "content": prompt}],
        )
    except Exception as e:
        return None, {"error": str(e), "input_tokens": 0, "output_tokens": 0}
    text = "".join(b.text for b in resp.content if hasattr(b, "text"))
    parsed = None
    try:
        if "```" in text:
            parts = text.split("```")
            for p in parts:
                p = p.strip()
                if p.startswith("json"):
                    p = p[4:].strip()
                if p.startswith("{"):
                    text = p; break
        first = text.find("{"); last = text.rfind("}")
        if first != -1 and last != -1:
            parsed = json.loads(text[first:last+1])
    except (json.JSONDecodeError, ValueError):
        parsed = None
    return parsed, {
        "input_tokens": resp.usage.input_tokens,
        "output_tokens": resp.usage.output_tokens,
    }


def haiku_cost(in_tok: int, out_tok: int) -> float:
    return (in_tok / 1_000_000) * 1.0 + (out_tok / 1_000_000) * 5.0


def main(argv: list[str]) -> int:
    p = argparse.ArgumentParser()
    p.add_argument("--dataset", required=True)
    p.add_argument("--judge-prompt", required=True)
    p.add_argument("--model", default=_get_default_model())
    p.add_argument("--max-cases", type=int, default=100)
    p.add_argument("--dry-run", action="store_true")
    args = p.parse_args(argv)

    dataset_path = Path(args.dataset)
    prompt_path = Path(args.judge_prompt)
    if not dataset_path.exists() or not prompt_path.exists():
        print(f"Missing: {dataset_path} or {prompt_path}", file=sys.stderr)
        return 2

    dataset = load_yaml(dataset_path)
    template = prompt_path.read_text()
    cases = dataset.get("cases", [])[: args.max_cases]
    if not cases:
        print("No cases."); return 0

    api_key = os.environ.get("ANTHROPIC_API_KEY")
    if not api_key and not args.dry_run:
        print("ANTHROPIC_API_KEY not set. Use --dry-run to preview prompts.", file=sys.stderr)
        return 2
    if not args.dry_run:
        from anthropic import Anthropic
        client = Anthropic(api_key=api_key)
    else:
        client = None

    runs_dir = dataset_path.parent / "runs"
    runs_dir.mkdir(exist_ok=True)
    run_id = datetime.now().strftime("%Y-%m-%d-%H%M%S")
    run_path = runs_dir / f"{run_id}.json"

    results: list[dict] = []
    tp = tn = fp = fn = 0
    total_cost = 0.0

    for case in cases:
        prompt = render_prompt(template, case)
        if args.dry_run:
            print(f"--- case {case['id']} ---")
            print(prompt[:500])
            print("...")
            results.append({"case_id": case["id"], "dry_run": True})
            continue
        parsed, usage = call_judge(client, args.model, prompt)
        cost = haiku_cost(usage["input_tokens"], usage["output_tokens"])
        total_cost += cost
        expected_verdict = case.get("expected", {}).get("verdict", "pass")
        actual_verdict = (parsed or {}).get("verdict", "error")
        match = (actual_verdict == expected_verdict)

        if expected_verdict == "fail":
            if actual_verdict == "fail":
                tp += 1
            else:
                fn += 1
        else:
            if actual_verdict == "pass":
                tn += 1
            elif actual_verdict == "fail":
                fp += 1

        results.append({
            "case_id": case["id"],
            "expected_verdict": expected_verdict,
            "actual_verdict": actual_verdict,
            "match": match,
            "scores": (parsed or {}).get("scores"),
            "reason": (parsed or {}).get("reason"),
            "confidence": (parsed or {}).get("confidence"),
            "cost_usd": cost,
            "usage": usage,
        })
        time.sleep(0.2)

    precision = tp / (tp + fp) if (tp + fp) else None
    recall = tp / (tp + fn) if (tp + fn) else None

    summary = {
        "run_id": run_id,
        "model": args.model,
        "dataset": str(dataset_path),
        "cases_total": len(cases),
        "tp": tp, "tn": tn, "fp": fp, "fn": fn,
        "precision": precision,
        "recall": recall,
        "cost_usd": total_cost,
        "results": results,
        "timestamp": datetime.now(timezone.utc).isoformat(),
    }
    run_path.write_text(json.dumps(summary, indent=2, ensure_ascii=False))
    print(f"[aria-eval-judge] run {run_id}: TP={tp} TN={tn} FP={fp} FN={fn} "
          f"precision={precision} recall={recall} cost=${total_cost:.4f}")
    print(f"  saved to {run_path}")
    return 0 if (fp + fn) == 0 else 1


if __name__ == "__main__":
    sys.exit(main(sys.argv[1:]))
