"""Conservative fixed-H minimax diagnostic from the validated p33 PASS."""
from __future__ import annotations

import argparse
from dataclasses import replace
import json
from pathlib import Path

import numpy as np

from ap1_r2c_implicit_dae_multiresolution import _physical_diagnostics
from ap1_r2c_implicit_dae_solver import ImplicitDAEResidual
from p42_fixed_h_minimax import solve_fixed_h_minimax


ROOT = Path(__file__).resolve().parent
if not (ROOT / "Apeiron").is_dir():
    ROOT = ROOT.parents[2]
STATE = ROOT / "Apeiron/HANDOFF/CURRENT_STATE_LATEST.npz"
AUDIT = ROOT / "Apeiron/AP1/APEIRON_AP1_R2C_G22_17_33_65_AUDIT_LATEST.json"


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

    audit = json.loads(AUDIT.read_text(encoding="utf-8"))
    level17 = next(level for level in audit["levels"] if level["candidate_nodes"] == 17)
    level33 = next(level for level in audit["levels"] if level["candidate_nodes"] == 33)
    if not level33.get("accepted_as_checkpoint"):
        raise RuntimeError("validated p33 PASS required")

    vector = np.asarray(level33["solution_vector"], dtype=float)
    system = ImplicitDAEResidual(STATE, nodes=33, uv_homotopy=1.0)
    baseline = _physical_diagnostics(system, vector)
    result, selected = solve_fixed_h_minimax(
        system,
        vector,
        radii=(1.0e-8, 3.0e-8, 1.0e-7, 3.0e-7, 1.0e-6, 3.0e-6),
        full_radius_diagnostics=True,
    )
    diagnostics = result["repeated_full_diagnostics"][0]
    precision_scan = []
    for digits in (60, 70, 80, 100):
        system.cfg = replace(system.cfg, digits=digits)
        item = _physical_diagnostics(system, selected)
        precision_scan.append({
            "digits": digits,
            "dae_max_abs_scaled_residual": item[
                "dae_max_abs_scaled_residual"
            ],
            "ward_normalized": item["ward_normalized"],
            "validation_source_relative_change": item[
                "validation_source_relative_change"
            ],
            "validation_wronskian_relative_error": item[
                "validation_wronskian_relative_error"
            ],
        })
    system.cfg = replace(system.cfg, digits=60)
    targets = level17["diagnostics"]
    result.update({
        "schema": "apeiron-ap1-r2c-p33-monotonic-minimax-diagnostic-v1.0",
        "source": "validated_p33_PASS_only",
        "baseline_diagnostics": baseline,
        "selected_precision_scan": precision_scan,
        "monotonic_targets_from_p17": {
            "ward_normalized": targets["ward_normalized"],
            "validation_source_relative_change": targets[
                "validation_source_relative_change"
            ],
            "validation_wronskian_relative_error": targets[
                "validation_wronskian_relative_error"
            ],
        },
        "monotonic_pass_against_p17": {
            "ward_normalized": diagnostics["ward_normalized"]
            <= targets["ward_normalized"],
            "validation_source_relative_change": diagnostics[
                "validation_source_relative_change"
            ] <= targets["validation_source_relative_change"],
            "validation_wronskian_relative_error": diagnostics[
                "validation_wronskian_relative_error"
            ] <= targets["validation_wronskian_relative_error"],
        },
        "candidate_retained": bool(
            result["accepted_as_checkpoint"]
            and diagnostics["ward_normalized"] <= targets["ward_normalized"]
            and diagnostics["validation_source_relative_change"]
            <= targets["validation_source_relative_change"]
            and diagnostics["validation_wronskian_relative_error"]
            <= targets["validation_wronskian_relative_error"]
        ),
    })
    if not result["candidate_retained"]:
        result["solution_vector"] = None
        result["trial_vector_stored_or_used_as_seed"] = False
    args.output.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8")


if __name__ == "__main__":
    main()
