#!/usr/bin/env python3
"""Verify published aggregates only. No private router, snapshot or inference.

The official corpus must be downloaded separately. This script never downloads
anything, executes upstream code, or reconstructs routing decisions.
"""
from __future__ import annotations

import argparse
from collections import Counter
from decimal import Decimal, localcontext
import hashlib
import importlib
import importlib.metadata
import json
import math
import numbers
import os
from pathlib import Path
import pickle
import sys

sys.dont_write_bytecode = True
ROWS = 36497
DATASET_BYTES = 99567659
DATASET_SHA256 = "ba4f77f19517610a707c374e99322d7750c30fc4ae7ff5527888595a1e65d36d"
PREDICTIONS_SHA256 = "d0c9e66e935a7fee78942cca78e16d9a10cd7a907b0f23ad9ee1b0a84d741155"
SUMMARY_SHA256 = "c2c4c9e427bc5369f5bcd7f1acdcfc921c41ddefce67f612332ba06df9a69ece"
ABS_TOLERANCE = 1e-12
REL_TOLERANCE = 1e-12
MODELS = (
    "gpt-3.5-turbo-1106", "claude-instant-v1", "claude-v1", "claude-v2",
    "gpt-4-1106-preview", "meta/llama-2-70b-chat", "mistralai/mixtral-8x7b-chat",
    "zero-one-ai/Yi-34B-Chat", "WizardLM/WizardLM-13B-V1.2",
    "meta/code-llama-instruct-34b-chat", "mistralai/mistral-7b-chat",
)
ALLOWED_GLOBALS = frozenset({
    ("numpy", "dtype"), ("numpy", "ndarray"),
    ("numpy.core.numeric", "_frombuffer"), ("numpy.core.multiarray", "_reconstruct"),
    ("pandas.core.frame", "DataFrame"), ("pandas.core.internals.managers", "BlockManager"),
    ("pandas._libs.internals", "_unpickle_block"),
    ("pandas.core.indexes.base", "_new_Index"), ("pandas.core.indexes.base", "Index"),
    ("pandas.core.indexes.range", "RangeIndex"), ("builtins", "slice"),
})
PREDICTION_FIELDS = {"sample_id", "model_id", "prompt_sha256", "latency_ms", "feature_sha256", "state_sha256", "formula_version", "eligible_count"}


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


def deny_network_and_execution(event, _args):
    if event in {"socket.connect", "socket.getaddrinfo", "socket.bind", "socket.sendto", "subprocess.Popen", "os.system", "os.fork"}:
        raise RuntimeError("Network and process execution are forbidden in this verifier")


class CorpusUnpickler(pickle.Unpickler):
    def find_class(self, module, name):
        if (module, name) not in ALLOWED_GLOBALS:
            raise pickle.UnpicklingError(f"Pickle global not allowed: {module}.{name}")
        return getattr(importlib.import_module(module), name)

    def persistent_load(self, _pid):
        raise pickle.UnpicklingError("Persistent pickle references are forbidden")


def reject_json_constant(_value):
    raise ValueError("Non-finite JSON values are forbidden")


def checked_bytes(path, expected_hash):
    raw = Path(path).read_bytes()
    require(hashlib.sha256(raw).hexdigest() == expected_hash, f"SHA-256 mismatch: {Path(path).name}")
    return raw


def load_corpus(path):
    import pandas as pd
    with Path(path).open("rb") as handle:
        require(os.fstat(handle.fileno()).st_size == DATASET_BYTES, "Official corpus size mismatch")
        require(hashlib.file_digest(handle, "sha256").hexdigest() == DATASET_SHA256, "Official corpus SHA-256 mismatch")
        handle.seek(0)
        frame = CorpusUnpickler(handle).load()
    require(type(frame) is pd.DataFrame and len(frame) == ROWS, "Unexpected corpus type or row count")
    columns = {"sample_id", "prompt", "eval_name", "oracle_model_to_route_to"} | set(MODELS)
    columns |= {model + suffix for model in MODELS for suffix in ("|total_cost", "|model_response")}
    require(set(frame.columns) == columns and len(frame.columns) == len(columns), "Unexpected corpus columns")
    return frame


def measurement(value, quality):
    import numpy as np
    require(isinstance(value, (numbers.Real, np.bool_)), "Missing or nonnumeric cached measurement")
    require(quality or not isinstance(value, (bool, np.bool_)), "Boolean cached cost is invalid")
    number = float(value)
    require(math.isfinite(number) and number >= 0 and (not quality or number <= 1), "Cached measurement outside its permitted range")
    return Decimal(str(number))


def verify(dataset, predictions_path, summary_path):
    for package, version in {"numpy": "1.26.4", "pandas": "2.2.3"}.items():
        require(importlib.metadata.version(package) == version, f"Use the pinned version: {package}=={version}")
    predictions = [json.loads(line, parse_constant=reject_json_constant) for line in checked_bytes(predictions_path, PREDICTIONS_SHA256).decode("utf-8").splitlines()]
    summary = json.loads(checked_bytes(summary_path, SUMMARY_SHA256), parse_constant=reject_json_constant)
    require(len(predictions) == ROWS, "Prediction count mismatch")
    require(all(set(row) == PREDICTION_FIELDS for row in predictions), "Unexpected public prediction fields")
    require(summary.get("rows") == ROWS and summary.get("excluded") == 0, "Summary count/exclusions mismatch")
    require(summary.get("dataset_sha256") == DATASET_SHA256 and summary.get("predictions_sha256") == PREDICTIONS_SHA256, "Summary provenance mismatch")
    require(summary.get("aiq") is None, "This verifier does not establish AIQ")
    frame = load_corpus(dataset)
    ids = frame["sample_id"].tolist()
    require(len(set(ids)) == ROWS and ids == [row.get("sample_id") for row in predictions], "Prediction IDs, order or uniqueness do not match the corpus")
    result_ids = ("titan-u-historical", *MODELS, "oracle-retrospective")
    totals = {name: {"quality": Decimal(0), "cost": Decimal(0)} for name in result_ids}
    selections = Counter({model: 0 for model in MODELS})
    columns = ["prompt", *MODELS, *[model + "|total_cost" for model in MODELS]]
    with localcontext() as context:
        context.prec = 50
        for prediction, record in zip(predictions, frame.loc[:, columns].itertuples(index=False, name=None)):
            prompt = record[0]
            require(isinstance(prompt, str) and hashlib.sha256(prompt.encode()).hexdigest() == prediction.get("prompt_sha256"), "Prediction prompt hash mismatch")
            selected_model = prediction.get("model_id")
            require(selected_model in MODELS, "Selection outside the frozen pool")
            quality = [measurement(value, True) for value in record[1:12]]
            cost = [measurement(value, False) for value in record[12:23]]
            best = min(range(11), key=lambda index: (-quality[index], cost[index], index))
            selection = {model: index for index, model in enumerate(MODELS)}
            selection["titan-u-historical"] = MODELS.index(selected_model)
            selection["oracle-retrospective"] = best
            for name, index in selection.items():
                totals[name]["quality"] += quality[index]
                totals[name]["cost"] += cost[index]
            selections[selected_model] += 1
        published_rows = summary.get("results")
        require(isinstance(published_rows, list) and len(published_rows) == len(result_ids), "Summary must contain all thirteen comparisons")
        published = {row["id"]: row for row in published_rows}
        require(set(published) == set(result_ids), "Summary comparison IDs mismatch")
        for name in result_ids:
            row = published[name]
            require(row.get("correct") is None and row.get("total") == ROWS, "Continuous quality must not be presented as binary accuracy")
            expected_kind = "router" if name == "titan-u-historical" else "oracle" if name == "oracle-retrospective" else "baseline"
            require(row.get("kind") == expected_kind, "Comparison role mismatch")
            recalculated_quality = float(totals[name]["quality"] / Decimal(ROWS))
            recalculated_cost = float(totals[name]["cost"])
            for key, recalculated in (("quality_score", recalculated_quality), ("estimated_corpus_cost_usd", recalculated_cost)):
                published_value = row.get(key)
                require(type(published_value) in (int, float) and math.isfinite(published_value), "Invalid published measurement")
                require(math.isclose(published_value, recalculated, rel_tol=REL_TOLERANCE, abs_tol=ABS_TOLERANCE), f"Aggregate mismatch: {name}/{key}")
            require(Decimal(row["cached_cost_usd_decimal"]) == totals[name]["cost"], f"Exact decimal cost mismatch: {name}")
    require(dict(selections) == summary.get("selection_counts"), "Published model selection counts mismatch")
    return {"verification": "PASS", "rows": ROWS, "comparisons": len(result_ids), "excluded": 0,
            "absolute_tolerance": ABS_TOLERANCE, "relative_tolerance": REL_TOLERANCE,
            "decimal_cost_check": "exact", "scope": "published aggregates only; not private selector reproduction"}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--dataset", type=Path, required=True)
    parser.add_argument("--predictions", type=Path, default=Path(__file__).with_name("predictions.jsonl"))
    parser.add_argument("--summary", type=Path, default=Path(__file__).with_name("summary.json"))
    args = parser.parse_args()
    sys.addaudithook(deny_network_and_execution)
    try:
        result = verify(args.dataset, args.predictions, args.summary)
    except Exception as error:
        print(json.dumps({"verification": "FAIL", "reason": str(error)}, ensure_ascii=False), file=sys.stderr)
        return 1
    print(json.dumps(result, ensure_ascii=False, sort_keys=True))
    return 0


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