#!/usr/bin/env python3
"""Read-only audit of the frozen pilot and its published results. No model calls.

Run after the pilot stops writing results:
  python3 scripts/verify-pilot.py --report /tmp/pilot-verification.json
Use --allow-incomplete during development; it permits missing planned episodes
and an absent download archive, but does not excuse malformed data or bad scores.
"""
import argparse
from collections import Counter
from datetime import datetime
import hashlib
import importlib.util
import json
from pathlib import Path, PurePosixPath
import re
import sys
import zipfile

sys.dont_write_bytecode = True
ROOT = Path(__file__).resolve().parents[1]
EXPERIMENT = ROOT / "experiments/company-evidence"
ARMS = ("notes", "beliefs", "dependencies")
SCORE_FIELDS = (
    "initial_decision_error", "post_revision_error", "recurrence_error",
    "later_decision_error", "copied_support_error", "independent_support_error",
    "support_count_errors", "unnecessary_abstentions", "unaffected_fact_errors",
)


def load(path):
    return json.loads(Path(path).read_text())


def compact(value):
    return json.dumps(value, ensure_ascii=False, separators=(",", ":"))


def size(value):
    return len(compact(value).encode("utf-8"))


def sha(path):
    return hashlib.sha256(Path(path).read_bytes()).hexdigest()


def require(condition, message):
    if not condition:
        raise AssertionError(message)


def import_runner(directory):
    spec = importlib.util.spec_from_file_location("frozen_evidence_pilot", directory / "run.py")
    runner = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(runner)
    return runner


def audit_fixtures(fixtures):
    require(len(fixtures) == 24, "Expected 24 fixtures: two development and four held-out families × four variants")
    require(len({fixture["id"] for fixture in fixtures}) == 24, "Fixture IDs must be unique")
    split_families = {}
    for split, expected_families in (("development", 2), ("heldout", 4)):
        group = [fixture for fixture in fixtures if fixture["split"] == split]
        families = sorted({fixture["family"] for fixture in group})
        require(len(families) == expected_families, f"Wrong {split} family count")
        for family in families:
            variants = Counter((f["initial_claim_true"], f["provenance"]) for f in group if f["family"] == family)
            require(variants == Counter({(truth, provenance): 1 for truth in (False, True) for provenance in ("copied", "independent")}), f"Unbalanced variants in {family}")
        split_families[split] = families
    require(not set(split_families["development"]) & set(split_families["heldout"]), "Development and held-out families overlap")
    for fixture in fixtures:
        stages = fixture["stages"]
        require(len(stages) == len(fixture["tasks"]) == len(fixture["expected"]) == 3, "Each fixture requires three stages")
        require([record["id"] for record in stages[0]] == ["R0", "R1", "R2", "R5"], "Initial record scope changed")
        require([record["id"] for record in stages[1]] == ["R3"], "Revision must arrive only at stage two")
        require([record["id"] for record in stages[2]] == ["R4"], "Republished note must arrive only at stage three")
        require(stages[0][0]["kind"] == "supplied_interpretation", "Inherited AI note must be explicitly supplied")
        for record in stages[0][1:3]:
            require(record["derived_from"] == (["R0"] if fixture["provenance"] == "copied" else []), "Source relationship disagrees with provenance condition")
            require(record["kind"] == ("internal_update" if fixture["provenance"] == "copied" else "direct_verification"), "Record kind disagrees with provenance condition")
        require(stages[1][0]["supersedes_sources"] == ["R0", "R1", "R2"], "Owner determination must supersede earlier sources")
        truth_word = "TRUE" if fixture["initial_claim_true"] else "FALSE"
        require(f"is {truth_word}." in stages[1][0]["text"], "Owner determination disagrees with fixture truth")
        require(stages[2][0]["derived_from"] == ["R0"], "Digest must trace back to the inherited note")
        expected = [
            {"decision": "hold", "support": 0} if fixture["provenance"] == "copied" else {"decision": "proceed", "support": 2},
            {"decision": "proceed", "support": 1} if fixture["initial_claim_true"] else {"decision": "do_not_proceed", "support": 0},
            {"decision": "proceed", "support": 1} if fixture["initial_claim_true"] else {"decision": "do_not_proceed", "support": 0},
        ]
        require(fixture["expected"] == expected, "Frozen scoring targets disagree with decision policy")
    return split_families


def score_episode(fixture, stages):
    """Recompute rubric outcomes from recorded model answers, including failures."""
    answers = [stage["parsed"] for stage in stages]
    wanted = fixture["expected"]
    wrong_decisions = [answer["decision"] != expected["decision"] for answer, expected in zip(answers, wanted)]
    support = [answer["independent_support_count"] for answer in answers]
    return {
        "initial_decision_error": wrong_decisions[0],
        "post_revision_error": wrong_decisions[1],
        "recurrence_error": bool(not fixture["initial_claim_true"] and answers[2]["decision"] == "proceed"),
        "later_decision_error": wrong_decisions[2],
        "copied_support_error": bool(fixture["provenance"] == "copied" and support[0] > 0),
        "independent_support_error": bool(fixture["provenance"] == "independent" and support[0] != 2),
        "support_count_errors": sum(count != expected["support"] for count, expected in zip(support, wanted)),
        "unnecessary_abstentions": sum(answer["decision"] == "hold" and expected["decision"] != "hold" for answer, expected in zip(answers, wanted)),
        "unaffected_fact_errors": sum(answer["handoff_deadline_status"] != "supported" for answer in answers),
    }


def audit_episode(path, fixture, runner):
    episode = load(path)
    label = episode["id"]
    arm = episode["arm"]
    require(arm in ARMS, f"{label}: unknown arm")
    require(label == f"{fixture['id']}-{arm}-{episode['repeat']}", f"{label}: ID disagrees with episode metadata")
    require(path.parent.name == label and path.parent.parent.name == fixture["split"], f"{label}: incorrect output path")
    for field in ("fixture_id", "family", "split", "initial_claim_true", "provenance"):
        wanted = fixture["id"] if field == "fixture_id" else fixture[field]
        require(episode[field] == wanted, f"{label}: wrong {field}")
    stages = episode["stages"]
    require(1 <= len(stages) <= 3, f"{label}: invalid call count")
    require([stage["stage"] for stage in stages] == list(range(1, len(stages) + 1)), f"{label}: missing or duplicate stage")
    raw_files = sorted(path.parent.glob("stage-*.json"))
    require([p.name for p in raw_files] == [f"stage-{n}.json" for n in range(1, len(stages) + 1)], f"{label}: raw call files disagree with episode")
    memory = "" if arm == "notes" else {"claims": [], "superseded_sources": [], "summary": ""}
    diagnostics = {"noncanonical_source_references": 0, "tool_events": 0, "stages": len(stages)}
    for index, stage in enumerate(stages):
        where = f"{label}/stage-{index + 1}"
        raw = load(raw_files[index])
        for key, value in raw["result"].items():
            require(stage[key] == value, f"{where}: episode disagrees with raw {key}")
        payload = json.loads(raw["input"].rsplit("\n\n", 1)[-1])
        require(set(payload) == {"proposition", "current_task", "stage", "records", "previous_memory"}, f"{where}: prompt includes fixture metadata or expected-answer keys")
        require(payload["proposition"] == fixture["proposition"], f"{where}: cross-fixture proposition")
        require(payload["stage"] == index + 1 and payload["current_task"] == fixture["tasks"][index], f"{where}: wrong current task")
        require(payload["records"] == fixture["stages"][index], f"{where}: current records differ (possible future/cross-split leakage)")
        require(stage["state_before_bytes"] == size(memory), f"{where}: state accounting mismatch before maintenance")
        if arm == "dependencies":
            carried, maintenance = runner.propagate(memory, fixture["stages"][index])
            require(stage["maintenance"] is not None, f"{where}: helper receipt missing")
            for key, value in maintenance.items():
                if key != "latency_ms":
                    require(stage["maintenance"][key] == value, f"{where}: helper receipt mismatch for {key}")
            require(stage["maintenance"]["latency_ms"] >= 0, f"{where}: negative helper duration")
        else:
            carried = memory
            require(stage["maintenance"] is None, f"{where}: unexpected helper in baseline")
        require(payload["previous_memory"] == carried, f"{where}: persistent input is not the preceding output plus declared maintenance")
        require(size(carried) <= runner.STATE_BYTES, f"{where}: maintained state exceeds common byte budget")
        prefix = runner.COMMON + "\n\n" + runner.ARMINSTRUCTION[arm]
        format_instruction = "\nFor structured memory, use at most 8 claims, keep eligibility.sources empty because its evidence comes through depends_on, and include handoff_deadline as a separate claim.\n\n"
        require(raw["input"] == prefix + format_instruction + compact(payload), f"{where}: prompt has unaccounted instructions or data")
        forbidden = ('"expected":', '"initial_claim_true":', '"split":', '"family":', '"provenance":')
        require(not any(key in compact({key: value for key, value in payload.items() if key != "previous_memory"}) for key in forbidden), f"{where}: answer metadata leaked")
        require(stage["elapsed_seconds"] >= 0, f"{where}: invalid duration")
        diagnostics["tool_events"] += len(stage.get("tool_events", []))
        if stage.get("error"):
            require(index == len(stages) - 1 and not episode["valid"], f"{where}: error must stop the episode")
            require(stage.get("parsed") is None and "state_after_bytes" not in stage, f"{where}: invalid output was persisted")
        else:
            parsed = json.loads(stage["output"])
            require(parsed == stage["parsed"], f"{where}: parsed output disagrees with raw output")
            runner.validate(parsed, arm)
            require(not stage.get("tool_events"), f"{where}: successful output used undeclared tools")
            require(stage.get("returncode") == 0, f"{where}: successful stage had a failed process")
            memory = parsed["memory"]
            require(stage["state_after_bytes"] == size(memory), f"{where}: output state byte count mismatch")
            if isinstance(memory, dict):
                refs = [ref for claim in memory["claims"] for ref in claim["sources"]] + memory["superseded_sources"]
                diagnostics["noncanonical_source_references"] += sum(re.fullmatch(r"R[0-9]+", ref) is None for ref in refs)
    valid = len(stages) == 3 and not any(stage.get("error") for stage in stages)
    require(episode["valid"] == valid, f"{label}: validity classification disagrees with calls")
    expected_scores = score_episode(fixture, stages) if valid else None
    require(episode["scores"] == expected_scores, f"{label}: saved scores disagree with raw answers")
    return episode, diagnostics


def aggregate(episodes):
    valid = [episode for episode in episodes if episode["valid"]]
    stages = [stage for episode in episodes for stage in episode["stages"]]
    state_sizes = [stage["state_after_bytes"] for stage in stages if "state_after_bytes" in stage]
    result = {
        "episodes": len(episodes), "valid_episodes": len(valid), "invalid_episodes": len(episodes) - len(valid),
        "copied_episodes": sum(episode["provenance"] == "copied" for episode in valid),
        "independent_episodes": sum(episode["provenance"] == "independent" for episode in valid),
        "false_claim_episodes": sum(not episode["initial_claim_true"] for episode in valid),
        "calls": len(stages), "call_wall_seconds": round(sum(stage["elapsed_seconds"] for stage in stages), 3),
        "state_observations": len(state_sizes), "mean_state_bytes": round(sum(state_sizes) / len(state_sizes), 1) if state_sizes else None,
        "helper_direct_marks": sum(len((stage.get("maintenance") or {}).get("directly_marked_claims", [])) for stage in stages),
        "helper_transitive_marks": sum(len((stage.get("maintenance") or {}).get("transitively_marked_claims", [])) for stage in stages),
        "helper_milliseconds": round(sum((stage.get("maintenance") or {}).get("latency_ms", 0) for stage in stages), 3),
    }
    result.update({field: sum(episode["scores"][field] for episode in valid) for field in SCORE_FIELDS})
    for field in ("input_tokens", "cached_input_tokens", "output_tokens", "reasoning_output_tokens"):
        result[field] = sum(stage["usage"][field] for stage in stages) if stages and all(field in stage.get("usage", {}) for stage in stages) else None
    return result


def audit_archive(path, experiment, results, primary_files):
    required = {name: experiment / name for name in primary_files}
    required.update({f"results/{name}": results / name for name in ("freeze.json", "summary.json")})
    for source in sorted((results / "runs").glob("*/*/*.json")):
        required[f"results/{source.relative_to(results)}"] = source
    with zipfile.ZipFile(path) as archive:
        names = archive.namelist()
        require(len(names) == len(set(names)), "Download archive has duplicate members")
        for name in names:
            parsed = PurePosixPath(name)
            require(not parsed.is_absolute() and ".." not in parsed.parts, "Download archive contains an unsafe path")
        for name, source in required.items():
            require(name in names, f"Download archive omits {name}")
            require(archive.read(name) == source.read_bytes(), f"Download archive differs from verified source: {name}")
    return {"members": len(names), "verified_primary_members": len(required), "sha256": sha(path)}


def verify(experiment, results, package, allow_incomplete=False):
    frozen = load(results / "freeze.json")
    primary_files = {"run.py", "fixtures.json", "PROTOCOL.md", *(f"schema-{arm}.json" for arm in ARMS)}
    require(set(frozen["files"]) == primary_files, "Freeze manifest must cover runner, protocol, fixtures and all three schemas")
    for name, expected in frozen["files"].items():
        require(sha(experiment / name) == expected, f"Frozen file changed: {name}")
    require(frozen["heldout_episodes"] == 144 and frozen["stages_per_episode"] == 3 and frozen["repetitions"] == 3, "Frozen run budget changed")
    require(frozen["maximum_model_calls_per_episode"] == 3 and frozen["state_limit_bytes"] == 5000, "Frozen resource allowance changed")
    runner = import_runner(experiment)
    fixtures = load(experiment / "fixtures.json")
    families = audit_fixtures(fixtures)
    require(fixtures == runner.make_fixtures(), "Generated fixtures differ from frozen fixture file")
    for arm in ARMS:
        require(load(experiment / f"schema-{arm}.json") == runner.schema(arm), f"Frozen schema disagrees with runner: {arm}")
    fixture_by_id = {fixture["id"]: fixture for fixture in fixtures}
    expected = {split: {f"{f['id']}-{arm}-{repeat}" for f in fixtures if f["split"] == split for arm in ARMS for repeat in range(1, 4 if split == "heldout" else 2)} for split in ("heldout", "development")}
    episodes, diagnostics = [], Counter()
    for path in sorted((results / "runs").glob("*/*/episode.json")):
        record = load(path)
        require(record["fixture_id"] in fixture_by_id, f"Unknown fixture in {path}")
        episode, observed = audit_episode(path, fixture_by_id[record["fixture_id"]], runner)
        require(episode["id"] in expected[episode["split"]], f"Unexpected run ID: {episode['id']}")
        if episode["split"] == "heldout":
            require(datetime.fromisoformat(episode["stages"][0]["started_at"]) >= datetime.fromisoformat(frozen["frozen_at"]), "Held-out run predates protocol freeze")
        episodes.append(episode)
        diagnostics.update(observed)
    require(len(episodes) == len({episode["id"] for episode in episodes}), "Duplicate episodes")
    actual = {split: {episode["id"] for episode in episodes if episode["split"] == split} for split in expected}
    missing = {split: sorted(expected[split] - actual[split]) for split in expected}
    if not allow_incomplete:
        require(not missing["heldout"], f"Held-out results incomplete: {len(actual['heldout'])}/144 episodes")
        require(not missing["development"], f"Development results incomplete: {len(actual['development'])}/24 episodes")
    summary = load(results / "summary.json")
    require(summary["heldout_target"] == 144 and summary["model"] == frozen["model"], "Summary model or target disagrees with freeze")
    require(summary["episode_count"] == len(episodes), "Summary is not synchronized with completed episode files; rerun after writer settles")
    require(summary["complete"] == (actual["heldout"] == expected["heldout"]), "Summary completeness is incorrect")
    require(summary["missing_heldout_ids"] == missing["heldout"] and summary["unexpected_heldout_ids"] == [], "Summary run ID accounting differs")
    recomputed = {}
    for split, section in (("heldout", "arms"), ("development", "development")):
        recomputed[section] = {}
        for arm in ARMS:
            values = aggregate([episode for episode in episodes if episode["split"] == split and episode["arm"] == arm])
            for field, value in values.items():
                require(summary[section][arm][field] == value, f"Summary {section}/{arm}/{field} differs from raw records: {summary[section][arm][field]!r} vs {value!r}")
            require(summary[section][arm]["valid_episodes"] + summary[section][arm]["invalid_episodes"] == summary[section][arm]["episodes"], "Invalid episodes disappeared from the denominator")
            recomputed[section][arm] = values
    archive = audit_archive(package, experiment, results, primary_files) if package.exists() else None
    require(allow_incomplete or archive is not None, "Published results archive is missing")
    return {
        "pass": True, "complete": not missing["heldout"] and not missing["development"], "allow_incomplete": allow_incomplete,
        "frozen_files": frozen["files"], "families": families,
        "completed_heldout_episodes": len(actual["heldout"]), "completed_development_episodes": len(actual["development"]),
        "missing_episode_counts": {key: len(value) for key, value in missing.items()},
        "diagnostics": dict(diagnostics), "recomputed": recomputed, "archive": archive,
        "limitations": ["This audits recorded prompts, outputs, score arithmetic and integrity; it does not independently authenticate the model service or prove the experimental hypothesis."],
    }


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--experiment", type=Path, default=EXPERIMENT)
    parser.add_argument("--results", type=Path)
    parser.add_argument("--package", type=Path, default=ROOT / "public/experiments/company-evidence/company-evidence-pilot.zip")
    parser.add_argument("--allow-incomplete", action="store_true")
    parser.add_argument("--report", type=Path)
    args = parser.parse_args()
    try:
        result = verify(args.experiment.resolve(), (args.results or args.experiment / "results").resolve(), args.package.resolve(), args.allow_incomplete)
    except (AssertionError, ValueError, KeyError, TypeError, OSError, zipfile.BadZipFile) as error:
        result = {"pass": False, "error": str(error), "allow_incomplete": args.allow_incomplete}
    rendered = json.dumps(result, indent=2) + "\n"
    if args.report:
        args.report.parent.mkdir(parents=True, exist_ok=True)
        args.report.write_text(rendered)
    print(rendered, end="")
    return 0 if result["pass"] else 1


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