"""Find a full p33 PASS between validated p17 and p65 diagnostic values."""
from __future__ import annotations

import argparse
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"
P65 = ROOT / "Apeiron/AP1/APEIRON_AP1_R2C_P65_RHO_LEXICOGRAPHIC_MINIMAX_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"))
    p17 = next(level for level in audit["levels"] if level["candidate_nodes"] == 17)
    p33 = next(level for level in audit["levels"] if level["candidate_nodes"] == 33)
    if not p33.get("accepted_as_checkpoint"):
        raise RuntimeError("validated p33 PASS required")
    p65 = json.loads(P65.read_text(encoding="utf-8"))
    if not p65.get("accepted_as_checkpoint"):
        raise RuntimeError("validated lexicographic p65 PASS required")
    p65_diagnostics = p65["repeated_full_diagnostics"][0]

    baseline = np.asarray(p33["solution_vector"], dtype=float)
    system = ImplicitDAEResidual(STATE, nodes=33, uv_homotopy=1.0)
    result, _selected = solve_fixed_h_minimax(
        system,
        baseline,
        radii=(1.0e-8, 3.0e-8, 1.0e-7, 3.0e-7, 1.0e-6, 3.0e-6, 1.0e-5),
        full_radius_diagnostics=True,
        objective_mode="rho_lexicographic",
        background_residual_cap=8.0e-5,
    )
    upper = p17["diagnostics"]
    lower = p65_diagnostics
    eligible = []
    compact = []
    for item in result["full_radius_diagnostics"]:
        diagnostics = item["full_diagnostics"]
        between = {
            key: lower[key] <= diagnostics[key] <= upper[key]
            for key in (
                "ward_normalized",
                "validation_source_relative_change",
            )
        }
        compact.append({
            "trust_radius": item["trust_radius"],
            "verified_background_max_abs": item[
                "verified_background_max_abs"
            ],
            "verified_rho_closure_max_abs": item[
                "verified_rho_closure_max_abs"
            ],
            "full_diagnostics": diagnostics,
            "all_absolute_gates_pass": item["all_absolute_gates_pass"],
            "between_p17_and_p65": between,
        })
        if item["all_absolute_gates_pass"] and all(between.values()):
            eligible.append((
                max(
                    (upper["ward_normalized"] - diagnostics["ward_normalized"])
                    / upper["ward_normalized"],
                    (diagnostics["validation_source_relative_change"]
                     - lower["validation_source_relative_change"])
                    / upper["validation_source_relative_change"],
                ),
                np.asarray(item["solution_vector"], dtype=float),
                compact[-1],
            ))
    selected = min(eligible, key=lambda row: row[0]) if eligible else None
    repeats = []
    if selected:
        repeats = [_physical_diagnostics(system, selected[1]) for _ in range(3)]
    payload = {
        "schema": "apeiron-ap1-r2c-p33-between-lexicographic-minimax-v1.0",
        "source": "validated_p33_PASS_only",
        "method": "fixed_H_velocity_minimax_with_background_cap_and_Q_rho_objective",
        "bounds": {
            "p17_upper": {
                "ward_normalized": upper["ward_normalized"],
                "validation_source_relative_change": upper[
                    "validation_source_relative_change"
                ],
            },
            "p65_lower": {
                "ward_normalized": lower["ward_normalized"],
                "validation_source_relative_change": lower[
                    "validation_source_relative_change"
                ],
            },
        },
        "radius_scan": compact,
        "eligible_count": len(eligible),
        "selected": selected[2] if selected else None,
        "repeated_full_diagnostics": repeats,
        "all_repeats_identical": bool(
            repeats and all(item == repeats[0] for item in repeats[1:])
        ),
        "accepted_as_checkpoint": selected is not None,
        "solution_vector": selected[1].tolist() if selected else None,
        "trial_vectors_stored_or_used_as_seed": False,
        "equations_changed": False,
        "physics_changed": False,
        "gates_changed": False,
        "production_curves_released": False,
    }
    args.output.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")


if __name__ == "__main__":
    main()
