"""Verify the structure and arithmetic of the stored 100-ID GMS evidence.

This verifier uses only Python's standard library and never starts the numerical
experiment, contacts a network service, opens a source checkout, or writes an
output file. It checks the immutable JSON record's declared configuration,
ordered records, and reported descriptive aggregates.

It cannot establish that the recorded source hashes match a checkout, that the
stored values came from a particular execution, or that the paper-level claims
hold. Those are intentionally outside this file-only check.
"""

from __future__ import annotations

import argparse
import hashlib
import json
import math
from pathlib import Path
import statistics
import sys
from typing import Any, Iterable, NoReturn


EVIDENCE_DIRECTORY = Path(__file__).resolve().parent
DEFAULT_INPUT = EVIDENCE_DIRECTORY / "gms-selected-ids-0-99.json"
METHODS = ("diffusion_resampling", "multinomial_baseline")
METRICS = ("sliced_wasserstein_l1", "posterior_mean_residual_squared_l2")
SUMMARY_FIELDS = ("mean", "population_standard_deviation", "minimum", "maximum")
EXPECTED_CONFIGURATION = {
    "author_keys_path": "experiments/rnd_keys.npy",
    "components": 5,
    "diffusion_a": -1.0,
    "diffusion_steps": 128,
    "dimension": 8,
    "execution": "one CPU process; compiled functions shared across ids",
    "integrator": "jentzen_and_kloeden",
    "mc_ids": list(range(100)),
    "observation_dimension": 1,
    "ode": True,
    "particles": 10_000,
    "swd_projections": 1_000,
    "terminal_time": 3.0,
}
EXPECTED_SOURCE_COMMIT = "767effe3e755067eb8a04422597fbf37eb8ab754"
EXPECTED_TOP_LEVEL_KEYS = {
    "aggregate",
    "claim_boundary",
    "configuration",
    "execution",
    "label",
    "per_id_results",
    "protocol_note",
    "reproduction_compatible",
    "source_commit",
    "source_paths",
    "source_provenance",
    "study_runtime_sha256",
}
ABSOLUTE_TOLERANCE = 1e-12


class VerificationError(ValueError):
    """Raised when a stored evidence record is malformed or internally inconsistent."""


def fail(message: str) -> NoReturn:
    raise VerificationError(message)


def sha256_file(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for block in iter(lambda: handle.read(1024 * 1024), b""):
            digest.update(block)
    return digest.hexdigest()


def strict_json_load(path: Path) -> dict[str, Any]:
    """Load a JSON object while rejecting duplicate keys and non-finite constants."""

    def reject_nonfinite(token: str) -> NoReturn:
        fail(f"{path.name} contains non-standard JSON constant {token!r}")

    def reject_duplicate_keys(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
        value: dict[str, Any] = {}
        for key, item in pairs:
            if key in value:
                fail(f"{path.name} contains duplicate JSON object key {key!r}")
            value[key] = item
        return value

    try:
        loaded = json.loads(
            path.read_text(encoding="utf-8"),
            parse_constant=reject_nonfinite,
            object_pairs_hook=reject_duplicate_keys,
        )
    except (OSError, json.JSONDecodeError) as error:
        fail(f"Cannot parse strict JSON from {path}: {error}")
    if not isinstance(loaded, dict):
        fail(f"Top-level value in {path.name} must be an object")
    return loaded


def require_object(value: object, path: str, *, keys: set[str] | None = None) -> dict[str, Any]:
    if not isinstance(value, dict):
        fail(f"{path} must be an object")
    if keys is not None and set(value) != keys:
        fail(f"{path} has keys {sorted(value)!r}; expected {sorted(keys)!r}")
    return value


def require_list(value: object, path: str) -> list[Any]:
    if not isinstance(value, list):
        fail(f"{path} must be a list")
    return value


def require_integer(value: object, path: str) -> int:
    if isinstance(value, bool) or not isinstance(value, int):
        fail(f"{path} must be an integer")
    return value


def require_number(value: object, path: str) -> float:
    if isinstance(value, bool) or not isinstance(value, (int, float)):
        fail(f"{path} must be a numeric JSON value")
    number = float(value)
    if not math.isfinite(number):
        fail(f"{path} must be finite")
    return number


def require_equal(actual: object, expected: object, path: str) -> None:
    if actual != expected:
        fail(f"{path} is {actual!r}; expected {expected!r}")


def require_close(actual: object, expected: float, path: str) -> None:
    value = require_number(actual, path)
    if not math.isclose(value, expected, rel_tol=0.0, abs_tol=ABSOLUTE_TOLERANCE):
        fail(f"{path} is {value:.17g}; recomputation is {expected:.17g}")


def summary(values: Iterable[float]) -> dict[str, float]:
    series = list(values)
    if not series:
        fail("Cannot aggregate an empty metric series")
    return {
        "mean": statistics.fmean(series),
        "population_standard_deviation": statistics.pstdev(series),
        "minimum": min(series),
        "maximum": max(series),
    }


def verify_declared_configuration(result: dict[str, Any]) -> None:
    configuration = require_object(
        result.get("configuration"), "configuration", keys=set(EXPECTED_CONFIGURATION)
    )
    for name, expected in EXPECTED_CONFIGURATION.items():
        value = configuration[name]
        path = f"configuration.{name}"
        if name == "mc_ids":
            identifiers = require_list(value, path)
            require_equal(
                [require_integer(item, f"{path}[{index}]") for index, item in enumerate(identifiers)],
                expected,
                path,
            )
            continue
        if isinstance(expected, bool):
            if not isinstance(value, bool):
                fail(f"{path} must be a Boolean")
        elif isinstance(expected, int):
            require_integer(value, path)
        elif isinstance(expected, float):
            require_number(value, path)
        require_equal(value, expected, path)

    require_equal(result.get("source_commit"), EXPECTED_SOURCE_COMMIT, "source_commit declaration")


def verify_records_and_aggregates(result: dict[str, Any]) -> dict[str, Any]:
    records = require_list(result.get("per_id_results"), "per_id_results")
    if len(records) != 100:
        fail(f"per_id_results has {len(records)} records; expected 100")

    series: dict[str, dict[str, list[float]]] = {
        method: {metric: [] for metric in METRICS} for method in METHODS
    }
    for expected_id, record_value in enumerate(records):
        record = require_object(
            record_value,
            f"per_id_results[{expected_id}]",
            keys={"mc_id", *METHODS},
        )
        require_equal(
            require_integer(record["mc_id"], f"per_id_results[{expected_id}].mc_id"),
            expected_id,
            f"per_id_results[{expected_id}].mc_id",
        )
        for method in METHODS:
            values = require_object(
                record[method],
                f"per_id_results[{expected_id}].{method}",
                keys=set(METRICS),
            )
            for metric in METRICS:
                series[method][metric].append(
                    require_number(
                        values[metric], f"per_id_results[{expected_id}].{method}.{metric}"
                    )
                )

    aggregate = require_object(
        result.get("aggregate"),
        "aggregate",
        keys={*METHODS, "paired_diffusion_minus_multinomial", "diffusion_lower_is_better_wins"},
    )
    recomputed: dict[str, Any] = {}
    for method in METHODS:
        reported_method = require_object(
            aggregate[method], f"aggregate.{method}", keys=set(METRICS)
        )
        recomputed[method] = {}
        for metric in METRICS:
            reported_summary = require_object(
                reported_method[metric],
                f"aggregate.{method}.{metric}",
                keys=set(SUMMARY_FIELDS),
            )
            calculated = summary(series[method][metric])
            for field, value in calculated.items():
                require_close(
                    reported_summary[field], value, f"aggregate.{method}.{metric}.{field}"
                )
            recomputed[method][metric] = calculated

    paired_reported = require_object(
        aggregate["paired_diffusion_minus_multinomial"],
        "aggregate.paired_diffusion_minus_multinomial",
        keys=set(METRICS),
    )
    wins_reported = require_object(
        aggregate["diffusion_lower_is_better_wins"],
        "aggregate.diffusion_lower_is_better_wins",
        keys=set(METRICS),
    )
    paired: dict[str, Any] = {}
    lower_counts: dict[str, int] = {}
    for metric in METRICS:
        differences = [
            diffusion - multinomial
            for diffusion, multinomial in zip(
                series["diffusion_resampling"][metric],
                series["multinomial_baseline"][metric],
                strict=True,
            )
        ]
        reported_summary = require_object(
            paired_reported[metric],
            f"aggregate.paired_diffusion_minus_multinomial.{metric}",
            keys=set(SUMMARY_FIELDS),
        )
        calculated = summary(differences)
        for field, value in calculated.items():
            require_close(
                reported_summary[field], value,
                f"aggregate.paired_diffusion_minus_multinomial.{metric}.{field}",
            )
        lower_count = sum(value < 0.0 for value in differences)
        require_equal(
            require_integer(
                wins_reported[metric], f"aggregate.diffusion_lower_is_better_wins.{metric}"
            ),
            lower_count,
            f"aggregate.diffusion_lower_is_better_wins.{metric}",
        )
        paired[metric] = calculated
        lower_counts[metric] = lower_count

    return {
        "record_count": len(records),
        "recomputed_aggregate": recomputed,
        "recomputed_paired_aggregate": paired,
        "diffusion_lower_is_better_wins": lower_counts,
    }


def verify_evidence(input_path: Path) -> dict[str, Any]:
    """Verify one stored result record and return the recomputed evidence summary."""

    result = strict_json_load(input_path)
    require_object(result, "top-level record", keys=EXPECTED_TOP_LEVEL_KEYS)
    verify_declared_configuration(result)
    recomputed = verify_records_and_aggregates(result)
    return {
        "input_sha256": sha256_file(input_path),
        "declared_source_commit": EXPECTED_SOURCE_COMMIT,
        **recomputed,
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--input",
        type=Path,
        default=DEFAULT_INPUT,
        help="stored GMS JSON to inspect (default: evidence/gms-selected-ids-0-99.json)",
    )
    return parser.parse_args()


def main() -> int:
    args = parse_args()
    try:
        report = verify_evidence(args.input)
    except (OSError, VerificationError) as error:
        print(f"Stored-evidence verification: failed: {error}", file=sys.stderr)
        return 1

    print("Stored-evidence verification: passed")
    print(f"Input SHA-256: {report['input_sha256']}")
    print(f"Ordered records checked: {report['record_count']} (IDs 0–99)")
    print("Methods checked: diffusion_resampling, multinomial_baseline")
    print("Metrics checked: sliced_wasserstein_l1, posterior_mean_residual_squared_l2")
    print(f"Diffusion-lower counts: {report['diffusion_lower_is_better_wins']}")
    print(
        "Recomputed from this file only: declared configuration, strict JSON structure, "
        "raw-record aggregates, paired aggregates, and lower-counts."
    )
    print(
        "Not established: source-checkout identity or hashes, numerical-execution provenance, "
        "paper reproduction, statistical significance, or any outperformance claim."
    )
    return 0


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