"""Bounded two-parent p-secant diagnostic for AP1-R2c.

Only complete accepted vectors from the atomic PASS-only checkpoint are read.
Every trial is checked on the full target DAE residual, and no trial vector is
written or reused as a continuation seed.
"""
from __future__ import annotations

import argparse
import json
from pathlib import Path

import numpy as np

from ap1_r2c_implicit_dae_multiresolution import (
    _lift_arrays,
    balanced_algebraic_projection,
)
from ap1_r2c_implicit_dae_solver import ImplicitDAEResidual


def _accepted_vector(report: dict, nodes: int) -> np.ndarray:
    levels = [
        level for level in report["levels"]
        if level.get("accepted_as_checkpoint")
        and level.get("candidate_nodes") == nodes
        and level.get("solution_vector")
    ]
    if not levels:
        raise ValueError(f"validated {nodes}-node parent required")
    return np.asarray(levels[-1]["solution_vector"], dtype=float)


def scan(state: Path, checkpoint: Path, target_nodes: int) -> dict:
    report = json.loads(checkpoint.read_text(encoding="utf-8"))
    immediate_nodes = target_nodes - 1
    older_nodes = target_nodes - 2
    immediate_vector = _accepted_vector(report, immediate_nodes)
    older_vector = _accepted_vector(report, older_nodes)
    immediate = ImplicitDAEResidual(
        state, nodes=immediate_nodes, uv_homotopy=1.0
    )
    older = ImplicitDAEResidual(state, nodes=older_nodes, uv_homotopy=1.0)
    target = ImplicitDAEResidual(state, nodes=target_nodes, uv_homotopy=1.0)
    trials = []

    for degree in (2, 3, 4, 5):
        newest = _lift_arrays(immediate, immediate_vector, target, "bspline", degree)
        previous = _lift_arrays(older, older_vector, target, "bspline", degree)
        for beta in np.linspace(-1.0, 3.0, 17):
            raw = newest + beta * (newest - previous)
            projected, projection = balanced_algebraic_projection(target, raw)
            norm = float(np.max(np.abs(target.evaluate(projected))))
            trials.append({
                "degree": degree,
                "beta": float(beta),
                "dae_max_abs_scaled_residual": norm,
                "projection_alpha": projection["selected_alpha"],
            })
    trials.sort(key=lambda item: item["dae_max_abs_scaled_residual"])
    return {
        "target_nodes": target_nodes,
        "parents": [older_nodes, immediate_nodes],
        "method": "same_degree_two_parent_p_secant",
        "beta_bounds": [-1.0, 3.0],
        "trial_vectors_stored_or_used_as_seed": False,
        "best": trials[:20],
    }


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("state", type=Path)
    parser.add_argument("checkpoint", type=Path)
    parser.add_argument("--target-nodes", type=int, required=True)
    args = parser.parse_args()
    print(json.dumps(scan(args.state, args.checkpoint, args.target_nodes), indent=2))


if __name__ == "__main__":
    main()
