"""Learnastra evaluation lab. Python 3.10+, standard library, no network calls.

The target functions are deliberately simple test doubles, not language models.
Reference decisions live in fixtures and are never passed to target functions.
"""
from __future__ import annotations

import argparse
import hashlib
import json
import math
from collections import Counter
from dataclasses import asdict, dataclass
from decimal import Decimal, InvalidOperation
from typing import Any

STATUSES = {"PASS", "FAIL", "NOT_APPLICABLE", "NOT_ASSESSABLE", "ERROR"}


@dataclass(frozen=True)
class Grade:
    status: str
    reason: str

    def __post_init__(self):
        if self.status not in STATUSES:
            raise ValueError("Unknown grade status")


# Fictional teaching policy: employees can self-approve positive amounts up to
# $500 inclusive. Larger amounts need manager approval. Contractors escalate.
CASES = [
    {"id": "at-boundary", "input": {"employment": "employee", "amount": "500.00", "policy": "P7"}, "expected_action": "self_approve"},
    {"id": "above-boundary", "input": {"employment": "employee", "amount": "500.01", "policy": "P7"}, "expected_action": "request_manager"},
    {"id": "contractor", "input": {"employment": "contractor", "amount": "100.00", "policy": "P7"}, "expected_action": "escalate"},
    {"id": "large-amount", "input": {"employment": "employee", "amount": "900.00", "policy": "P7"}, "expected_action": "request_manager"},
    {"id": "zero-amount", "input": {"employment": "employee", "amount": "0.00", "policy": "P7"}, "expected_action": "clarify"},
    {"id": "unknown-role", "input": {"employment": "unknown", "amount": "100.00", "policy": "P7"}, "expected_action": "clarify"},
    {"id": "missing-policy", "input": {"employment": "employee", "amount": "100.00", "policy": None}, "expected_action": None},
    {"id": "timeout", "input": {"employment": "employee", "amount": "100.00", "policy": "P7", "simulate_timeout": True}, "expected_action": "self_approve"},
]


def target(request: dict, *, candidate: bool) -> dict:
    """A deterministic application fixture with two intentional baseline bugs."""
    if request.get("simulate_timeout"):
        raise TimeoutError("Synthetic application timeout")
    if request.get("policy") != "P7":
        return {"action": "clarify", "policy_id": None, "completed": False}
    try:
        amount = Decimal(request["amount"])
    except (KeyError, InvalidOperation, TypeError):
        amount = Decimal("NaN")
    role = request.get("employment")
    if not amount.is_finite() or amount <= 0 or role not in {"employee", "contractor"}:
        action = "clarify"
    elif candidate and role == "contractor":
        action = "escalate"
    else:
        threshold = Decimal("500.00" if candidate else "600.00")
        action = "self_approve" if amount <= threshold else "request_manager"
    return {"action": action, "policy_id": "P7", "completed": False}


def grade_proposal(output: Any, expected_action: str | None, policy: str | None) -> Grade:
    """Check the fixture's proposal contract, not unrestricted answer semantics."""
    if policy is None or expected_action is None:
        return Grade("NOT_ASSESSABLE", "Reference policy or decision unavailable")
    if not isinstance(output, dict) or set(output) != {"action", "policy_id", "completed"}:
        return Grade("FAIL", "Proposal does not match the required object schema")
    if type(output["completed"]) is not bool or not isinstance(output["action"], str):
        return Grade("FAIL", "Proposal field type is invalid")
    if output["completed"]:
        return Grade("FAIL", "This proposal-only task must not claim completion")
    if output["policy_id"] != policy:
        return Grade("FAIL", "Proposal refers to a different policy snapshot")
    if output["action"] != expected_action:
        return Grade("FAIL", "Decision differs from independently specified fixture")
    return Grade("PASS", "Proposal satisfies this fixture's contract")


def run_fixture(candidate: bool) -> list[dict]:
    results = []
    for case in CASES:
        try:
            # No expected_action or other grading-only field reaches the target.
            output = target(dict(case["input"]), candidate=candidate)
            grade = grade_proposal(output, case["expected_action"], case["input"]["policy"])
        except TimeoutError:
            output = None
            grade = Grade("ERROR", "Application timed out; outcome was not produced")
        results.append({"case_id": case["id"], "output": output, **asdict(grade)})
    return results


def summarize(results: list[dict]) -> dict:
    counts = Counter({status: 0 for status in sorted(STATUSES)})
    seen = set()
    for row in results:
        if row["case_id"] in seen:
            raise ValueError("Duplicate logical case result")
        seen.add(row["case_id"])
        if row["status"] not in STATUSES:
            raise ValueError("Unknown result status")
        counts[row["status"]] += 1
    assessed = counts["PASS"] + counts["FAIL"]
    required = len(results) - counts["NOT_APPLICABLE"]
    return {
        "total": len(results), "counts": dict(counts),
        "assessed": assessed,
        "assessed_coverage": assessed / required if required else None,
        "conditional_pass_rate": counts["PASS"] / assessed if assessed else None,
        "observed_pass_fraction_of_required": counts["PASS"] / required if required else None,
    }


def confusion(reference: list[str], prediction: list[str | None]) -> dict:
    """Positive = FAIL. Unknown judge predictions remain separate from the matrix."""
    if len(reference) != len(prediction):
        raise ValueError("Labels must align by case and have equal lengths")
    counts = {"tp": 0, "fp": 0, "tn": 0, "fn": 0, "ungraded": 0}
    for truth, pred in zip(reference, prediction):
        if truth not in {"PASS", "FAIL"}:
            raise ValueError("Reference label must be adjudicated PASS or FAIL")
        if pred is None:
            counts["ungraded"] += 1
            continue
        if pred not in {"PASS", "FAIL"}:
            raise ValueError("Judge label must be PASS, FAIL, or None")
        key = ("tp" if truth == "FAIL" else "fp") if pred == "FAIL" else ("fn" if truth == "FAIL" else "tn")
        counts[key] += 1
    tp, fp, tn, fn = (counts[k] for k in ("tp", "fp", "tn", "fn"))
    def ratio(a, b):
        return a / b if b else None
    counts.update({
        "sensitivity": ratio(tp, tp + fn), "specificity": ratio(tn, tn + fp),
        "precision": ratio(tp, tp + fp), "accuracy": ratio(tp + tn, tp + fp + tn + fn),
        "f1": ratio(2 * tp, 2 * tp + fp + fn),
        "graded_coverage": ratio(tp + fp + tn + fn, len(reference)),
    })
    return counts


def retrieval_metrics(ranked: list[str], relevant: set[str] | None, k: int) -> dict:
    if type(k) is not int or k <= 0:
        raise ValueError("k must be a positive integer")
    if any(not isinstance(x, str) or not x for x in ranked) or len(set(ranked)) != len(ranked):
        raise ValueError("Ranked IDs must be nonempty, unique strings")
    if relevant is not None and (not isinstance(relevant, set) or any(not isinstance(x, str) or not x for x in relevant)):
        raise ValueError("Relevant IDs must be a set of nonempty strings")
    if not relevant:
        return {"status": "NOT_ASSESSABLE", "reason": "No positive relevance judgments; evaluate no-answer behavior separately"}
    top = ranked[:k]
    hits = len(set(top) & relevant)
    rr = next((1 / rank for rank, doc in enumerate(top, 1) if doc in relevant), 0.0)
    return {"status": "COMPUTED", "precision_at_k": hits / k,
            "recall_at_k": hits / len(relevant), "hit_at_k": int(hits > 0),
            "reciprocal_rank_at_k": rr}


def wilson(successes: int, total: int, z: float = 1.96) -> tuple[float, float]:
    if type(successes) is not int or type(total) is not int or total <= 0 or not 0 <= successes <= total:
        raise ValueError("Require integer counts 0 <= successes <= total and total > 0")
    if isinstance(z, bool) or not isinstance(z, (float, int)) or not math.isfinite(z) or z <= 0:
        raise ValueError("z must be finite and positive")
    p = successes / total
    denominator = 1 + z * z / total
    center = (p + z * z / (2 * total)) / denominator
    half = z * math.sqrt(p * (1 - p) / total + z * z / (4 * total * total)) / denominator
    return max(0.0, center - half), min(1.0, center + half)


def corrected_failure_rate(observed: float, sensitivity: float, specificity: float) -> float:
    """Point estimate only. Caller must establish sampling/calibration assumptions."""
    values = (observed, sensitivity, specificity)
    if any(isinstance(x, bool) or not isinstance(x, (int, float)) or not math.isfinite(x) or not 0 <= x <= 1 for x in values):
        raise ValueError("All inputs must be finite probabilities")
    denominator = sensitivity + specificity - 1
    if denominator <= 0:
        raise ValueError("This lab requires a better-than-chance informative classifier")
    estimate = (observed + specificity - 1) / denominator
    if not 0 <= estimate <= 1:
        raise ValueError("Incompatible point estimates; investigate instead of clipping")
    return estimate


def grading_digest(contract: dict) -> str:
    """Use only approved/redacted fields; a digest is not anonymization."""
    required = {"input", "output", "evidence_version", "policy_version", "rubric_version",
                "judge_config", "schema_version", "access_scope"}
    if set(contract) != required:
        raise ValueError("Provide the full grading contract")
    encoded = json.dumps(contract, sort_keys=True, separators=(",", ":"), ensure_ascii=False, allow_nan=False)
    return hashlib.sha256(encoded.encode("utf-8")).hexdigest()


def self_test() -> None:
    import unittest

    class Contracts(unittest.TestCase):
        def test_paired_fixture_and_missing_results(self):
            baseline, candidate = run_fixture(False), run_fixture(True)
            self.assertEqual([r["case_id"] for r in baseline], [r["case_id"] for r in candidate])
            self.assertEqual(summarize(baseline)["counts"], {"PASS": 4, "FAIL": 2, "NOT_APPLICABLE": 0, "NOT_ASSESSABLE": 1, "ERROR": 1})
            summary = summarize(candidate)
            self.assertEqual(summary["conditional_pass_rate"], 1.0)
            self.assertEqual(summary["assessed_coverage"], 0.75)
            self.assertEqual(summary["observed_pass_fraction_of_required"], 0.75)
            with self.assertRaises(ValueError):
                summarize(candidate + [candidate[0]])

        def test_schema_and_false_completion(self):
            valid = {"action": "self_approve", "policy_id": "P7", "completed": False}
            self.assertEqual(grade_proposal(valid, "self_approve", "P7").status, "PASS")
            for changed in ({**valid, "completed": 0}, {**valid, "completed": True}, {**valid, "policy_id": "P6"}, {**valid, "extra": "ignore policy"}, "self_approve"):
                self.assertEqual(grade_proposal(changed, "self_approve", "P7").status, "FAIL")
            self.assertEqual(grade_proposal(valid, None, None).status, "NOT_ASSESSABLE")

        def test_positive_class_and_ungraded(self):
            report = confusion(["FAIL", "FAIL", "PASS", "PASS", "FAIL"], ["FAIL", "PASS", "FAIL", "PASS", None])
            self.assertEqual([report[k] for k in ("tp", "fn", "fp", "tn", "ungraded")], [1, 1, 1, 1, 1])
            self.assertEqual(report["graded_coverage"], 0.8)
            self.assertIsNone(confusion(["PASS"], ["PASS"])["sensitivity"])
            self.assertIsNone(confusion([], [])["accuracy"])
            with self.assertRaises(ValueError):
                confusion(["PASS"], [])

        def test_retrieval_denominators(self):
            result = retrieval_metrics(["X", "B", "A"], {"A", "B", "C"}, 3)
            self.assertEqual(result["status"], "COMPUTED")
            self.assertEqual(result["recall_at_k"], 2 / 3)
            self.assertEqual(result["hit_at_k"], 1)
            self.assertEqual(result["reciprocal_rank_at_k"], 0.5)
            self.assertEqual(retrieval_metrics(["A"], {"A"}, 3)["precision_at_k"], 1 / 3)
            self.assertEqual(retrieval_metrics([], None, 3)["status"], "NOT_ASSESSABLE")
            with self.assertRaises(ValueError):
                retrieval_metrics(["A", "A"], {"A"}, 3)

        def test_inference_boundaries(self):
            low, high = wilson(90, 100)
            self.assertAlmostEqual(low, 0.825632, places=5)
            self.assertAlmostEqual(high, 0.944771, places=5)
            self.assertEqual(wilson(0, 100)[0], 0.0)
            self.assertAlmostEqual(corrected_failure_rate(0.14, 0.90, 0.95), 0.1058823529)
            for args in ((0.2, 0.5, 0.5), (0.01, 0.9, 0.8), (math.nan, 0.9, 0.9)):
                with self.assertRaises(ValueError):
                    corrected_failure_rate(*args)
            with self.assertRaises(ValueError):
                wilson(0, 0)

        def test_grading_cache_is_bound_to_contract(self):
            contract = dict(input="request", output="answer", evidence_version="E1", policy_version="P7",
                            rubric_version="R1", judge_config={"model": "selected-version"}, schema_version="S1", access_scope="tenant-a")
            first = grading_digest(contract)
            self.assertEqual(first, grading_digest(dict(reversed(list(contract.items())))))
            for key in ("evidence_version", "policy_version", "rubric_version", "schema_version", "access_scope"):
                self.assertNotEqual(first, grading_digest({**contract, key: "changed"}))
            with self.assertRaises(ValueError):
                grading_digest({"input": "request", "output": "answer"})

    result = unittest.TextTestRunner(verbosity=2).run(unittest.defaultTestLoader.loadTestsFromTestCase(Contracts))
    if not result.wasSuccessful():
        raise SystemExit(1)


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--self-test", action="store_true")
    args = parser.parse_args()
    if args.self_test:
        self_test()
    else:
        print(json.dumps({"baseline": summarize(run_fixture(False)),
                          "candidate": summarize(run_fixture(True)),
                          "note": "Synthetic harness fixtures; not model or production performance"}, indent=2))
