# Copied verbatim from `evals/harness/golden_harness.py` by `scripts/publish-process-artifacts.ts`.
# Do not edit here — edit the source and rebuild.

#!/usr/bin/env python3
"""Runs golden_set.json (50 Q/A pairs) against a deployed /api/research-chat and scores
answer/citation/groundedness/limitation-inclusion (in-corpus) and correct-refusal-rate
(out-of-scope), per the dimensions defined in evals/golden_set.md.

Usage:
    python -m harness.golden_harness --base-url http://localhost:3000
    python -m harness.golden_harness --base-url https://jasonstiltner.com --out results/golden_run.json

Each item gets a fresh session (new_session=True) so no item's conversation history can leak
into another's — the golden set is a set of independent single-turn probes, not a
conversation, and reusing a session would let an earlier item's context change how a later
one is answered.
"""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import asdict
from datetime import datetime, timezone
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from harness.endpoint_client import ResearchChatClient  # noqa: E402
from harness import judge  # noqa: E402

GOLDEN_SET_PATH = Path(__file__).resolve().parent.parent / "golden_set.json"


def run_item(client: ResearchChatClient, item: dict, judge_model: str) -> dict:
    result = client.send(item["question"], new_session=True)
    cited_sources = [
        {"index": s.index, "page_title": s.page_title, "section_title": s.section_title, "url": s.url}
        for s in result.sources
    ]
    record = {
        "id": item["id"],
        "category": item["category"],
        "question": item["question"],
        "answer_text": result.answer_text,
        "cited_sources": cited_sources,
        "model": result.model,
        "band": result.band,
        "top_score": result.top_score,
        "route_refused": result.refused,
        "refusal_reason": result.refusal_reason,
        "http_status": result.http_status,
    }

    if item["category"] == "out_of_scope":
        text_to_grade = result.refusal_message if result.refused else result.answer_text
        verdict = judge.score_refusal_item(
            question=item["question"],
            reference_behavior=item["reference_behavior"],
            generated_answer=text_to_grade or "(empty response)",
            model=judge_model,
        )
        record["verdict"] = asdict(verdict)
        record["pass"] = verdict.correctly_refused
    else:
        if result.refused:
            # An in-corpus item that got refused at the confidence gate is a real failure —
            # the retrieval threshold was miscalibrated for this question, not a scoring
            # ambiguity. Fail it directly without spending a judge call on an empty answer.
            record["verdict"] = None
            record["pass"] = False
            record["fail_reason"] = f"wrongly refused (band={result.band}, top_score={result.top_score})"
        else:
            verdict = judge.score_golden_item(
                question=item["question"],
                reference_answer=item["reference_answer"],
                requires_limitation=item["requires_limitation"],
                generated_answer=result.answer_text,
                cited_sources=cited_sources,
                model=judge_model,
            )
            record["verdict"] = asdict(verdict)
            dims = [verdict.answer_correct, verdict.citation_correct, verdict.grounded]
            if item["requires_limitation"]:
                dims.append(bool(verdict.limitation_included))
            record["pass"] = all(dims)

    return record


def summarize(records: list[dict]) -> dict:
    graded = [r for r in records if not r.get("judge_error")]
    judge_errors = [r["id"] for r in records if r.get("judge_error")]

    in_corpus = [r for r in graded if r["category"] != "out_of_scope"]
    out_of_scope = [r for r in graded if r["category"] == "out_of_scope"]
    limitation_items = [r for r in in_corpus if r["category"] == "limitation_inclusive"]

    def rate(items: list[dict], key) -> float | None:
        return round(sum(1 for i in items if key(i)) / len(items), 4) if items else None

    citation_correct_items = [r for r in in_corpus if r.get("verdict")]
    return {
        "total_items": len(records),
        "graded_items": len(graded),
        "judge_errors": judge_errors,
        "in_corpus_answer_correctness": rate(in_corpus, lambda r: r.get("verdict") and r["verdict"]["answer_correct"]),
        "citation_correctness": rate(citation_correct_items, lambda r: r["verdict"]["citation_correct"]),
        "groundedness": rate(citation_correct_items, lambda r: r["verdict"]["grounded"]),
        "limitation_inclusive_answer_correctness": rate(limitation_items, lambda r: r["pass"]),
        "correct_refusal_rate": rate(out_of_scope, lambda r: r["pass"]),
        "overall_pass_rate": rate(graded, lambda r: r["pass"]),
        "wrongly_refused_in_corpus": [r["id"] for r in in_corpus if r.get("fail_reason", "").startswith("wrongly refused")],
        "failed_items": [r["id"] for r in graded if not r["pass"]],
    }


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--base-url", required=True, help="e.g. http://localhost:3000 or https://jasonstiltner.com")
    parser.add_argument("--out", default=None, help="Where to write the results JSON (default: stdout only)")
    parser.add_argument("--min-interval-s", type=float, default=21.0, help="Minimum gap between requests. Voyage's free-tier cap (3 RPM, account-wide, shared across every concurrent caller) is the real constraint, not a nicety -- 21s keeps a single sequential caller safely under it, matching scripts/build-corpus-index.ts's proven-safe pacing. Confirmed live: 5s (12/min) was NOT safe and produced spurious 503s misread as quality failures.")
    parser.add_argument("--judge-model", default=judge.JUDGE_MODEL_DEFAULT)
    parser.add_argument("--limit", type=int, default=None, help="Only run the first N items (smoke testing)")
    parser.add_argument("--ids", nargs="*", default=None, help="Only run these specific item ids")
    args = parser.parse_args()

    golden = json.loads(GOLDEN_SET_PATH.read_text(encoding="utf-8"))
    items = golden["items"]
    if args.ids:
        wanted = set(args.ids)
        items = [i for i in items if i["id"] in wanted]
    if args.limit:
        items = items[: args.limit]

    client = ResearchChatClient(args.base_url, min_interval_s=args.min_interval_s)

    records = []
    for i, item in enumerate(items, 1):
        print(f"[{i}/{len(items)}] {item['id']}: {item['question'][:70]}", file=sys.stderr)
        try:
            record = run_item(client, item, args.judge_model)
        except (judge.JudgeRefusedError, judge.JudgeFormatError) as e:
            # The judge, not the endpoint, failed — not a security/quality fail of the system
            # under test. See both exceptions' docstrings.
            record = {"id": item["id"], "category": item["category"], "question": item["question"],
                       "pass": None, "judge_error": True, "error": str(e)}
        except Exception as e:  # noqa: BLE001 - a transient API failure shouldn't kill a 50-item run
            record = {"id": item["id"], "category": item["category"], "question": item["question"],
                       "pass": False, "error": f"{type(e).__name__}: {e}"}
        status = "JUDGE_ERROR" if record.get("judge_error") else ("PASS" if record["pass"] else "FAIL")
        print(f"    -> {status}" + (f" ({record['error']})" if "error" in record else ""), file=sys.stderr)
        records.append(record)

    summary = summarize(records)
    report = {
        "run_at": datetime.now(timezone.utc).isoformat(),
        "base_url": args.base_url,
        "judge_model": args.judge_model,
        "confirmed_thresholds": golden["_meta"]["confirmed_thresholds"],
        "summary": summary,
        "items": records,
    }

    output = json.dumps(report, indent=2)
    if args.out:
        Path(args.out).write_text(output, encoding="utf-8")
        print(f"\nWrote {args.out}", file=sys.stderr)
    print(json.dumps(summary, indent=2))

    thresholds = golden["_meta"]["confirmed_thresholds"]
    ok = True
    if summary["citation_correctness"] is not None and summary["citation_correctness"] < thresholds["citation_correctness_min"]:
        ok = False
    if summary["correct_refusal_rate"] is not None and summary["correct_refusal_rate"] < thresholds["correct_refusal_rate_min"]:
        ok = False
    if summary["limitation_inclusive_answer_correctness"] is not None and summary["limitation_inclusive_answer_correctness"] < thresholds["limitation_inclusive_answer_correctness_min"]:
        ok = False
    return 0 if ok else 1


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