"""Final frozen G22 audit with PASS-only states and precision refinement."""
from __future__ import annotations

import argparse
from dataclasses import replace
from datetime import datetime, timezone
import json
from pathlib import Path

import numpy as np

from ap1_r2c_implicit_dae_multiresolution import (
    DAE_RESIDUAL_GATE,
    G22_THRESHOLDS,
    _physical_diagnostics,
    classify,
)
from ap1_r2c_implicit_dae_solver import ImplicitDAEResidual


ROOT = Path(__file__).resolve().parent
if not (ROOT / "Apeiron").is_dir():
    ROOT = ROOT.parents[2]
STATE = ROOT / "Apeiron/HANDOFF/CURRENT_STATE_LATEST.npz"
OLD_AUDIT = ROOT / "Apeiron/AP1/APEIRON_AP1_R2C_G22_17_33_65_AUDIT_LATEST.json"
P33 = ROOT / "Apeiron/AP1/APEIRON_AP1_R2C_P33_BETWEEN_LEXICOGRAPHIC_MINIMAX_LATEST.json"
P65 = ROOT / "Apeiron/AP1/APEIRON_AP1_R2C_P65_RHO_LEXICOGRAPHIC_MINIMAX_LATEST.json"
PRECISION_SCHEDULE = {17: 60, 33: 70, 65: 80}


def repeated(system: ImplicitDAEResidual, vector: np.ndarray) -> dict:
    samples = [_physical_diagnostics(system, vector) for _ in range(3)]
    keys = (
        "dae_max_abs_scaled_residual",
        "max_abs_friedmann_residual",
        "ward_normalized",
        "validation_source_relative_change",
        "validation_wronskian_relative_error",
    )
    conservative = {key: max(sample[key] for sample in samples) for key in keys}
    conservative.update({
        "finite": all(sample["finite"] for sample in samples),
        "positive_H": all(sample["positive_H"] for sample in samples),
        "min_H": min(sample["min_H"] for sample in samples),
        "min_PX": min(sample["min_PX"] for sample in samples),
        "min_K": min(sample["min_K"] for sample in samples),
    })
    conservative["gate_pass"] = samples[0]["gate_pass"]
    return {
        "repeats": 3,
        "samples": samples,
        "all_repeats_identical": all(item == samples[0] for item in samples[1:]),
        "conservative": conservative,
    }


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--output", type=Path, required=True)
    args = parser.parse_args()

    old = json.loads(OLD_AUDIT.read_text(encoding="utf-8"))
    old17 = next(level for level in old["levels"] if level["candidate_nodes"] == 17)
    p33 = json.loads(P33.read_text(encoding="utf-8"))
    p65 = json.loads(P65.read_text(encoding="utf-8"))
    if not old17.get("accepted_as_checkpoint"):
        raise RuntimeError("validated p17 PASS required")
    if not p33.get("accepted_as_checkpoint"):
        raise RuntimeError("validated p33 PASS required")
    if not p65.get("accepted_as_checkpoint"):
        raise RuntimeError("validated p65 PASS required")

    vectors = {
        17: np.asarray(old17["solution_vector"], dtype=float),
        33: np.asarray(p33["solution_vector"], dtype=float),
        65: np.asarray(p65["solution_vector"], dtype=float),
    }
    sources = {
        17: "validated_lambda1_PASS",
        33: "validated_p33_lexicographic_PASS_only",
        65: "validated_p65_lexicographic_PASS_only",
    }
    levels = []
    for nodes in (17, 33, 65):
        system = ImplicitDAEResidual(STATE, nodes=nodes, uv_homotopy=1.0)
        system.cfg = replace(system.cfg, digits=PRECISION_SCHEDULE[nodes])
        validation = repeated(system, vectors[nodes])
        diagnostics = validation["conservative"]
        levels.append({
            "candidate_nodes": nodes,
            "decimal_mode_digits": PRECISION_SCHEDULE[nodes],
            "source": sources[nodes],
            "diagnostics": diagnostics,
            "repeated_validation": validation,
            "accepted_as_checkpoint": all(diagnostics["gate_pass"].values()),
            "solution_vector": vectors[nodes].tolist(),
        })

    verdict = classify(levels)
    payload = {
        "schema": "apeiron-ap1-r2c-g22-monotonic-precision-audit-v1.0",
        "updated_utc": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"),
        "classification": verdict["classification"],
        "thresholds_frozen": G22_THRESHOLDS,
        "dae_residual_gate": DAE_RESIDUAL_GATE,
        "numerical_method_extension": {
            "fixed_H_lexicographic_Q_rho_minimax": True,
            "background_residual_cap": 8.0e-5,
            "decimal_precision_schedule": PRECISION_SCHEDULE,
            "precision_rationale": (
                "Increase Decimal guard digits with p so accumulated symplectic "
                "roundoff is itself refined; Float64 background inputs, equations, "
                "physical parameters, and all gate thresholds remain unchanged."
            ),
        },
        "levels": levels,
        "verdict": verdict,
        "absolute_gates_all_levels_pass": all(
            level["accepted_as_checkpoint"] for level in levels
        ),
        "equations_physics_or_gates_changed": False,
        "nonpass_stored_or_used_as_seed": False,
        "production_curves_released": False,
    }
    args.output.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")


if __name__ == "__main__":
    main()
