"""Bounded diagnostic scan for one AP1-R2c p-continuation level.

The script only reads a validated parent from the pass-only checkpoint.  Every
trial is independently algebraically projected and evaluated against the full
unchanged DAE residual.  Trial vectors are never written or reused as seeds.
"""
from __future__ import annotations

import argparse
import json
from pathlib import Path

import numpy as np

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


def scan(state: Path, checkpoint: Path, target_nodes: int) -> dict:
    report = json.loads(checkpoint.read_text(encoding="utf-8"))
    parents = [
        level for level in report["levels"]
        if level.get("accepted_as_checkpoint")
        and level.get("candidate_nodes") == target_nodes - 1
        and level.get("solution_vector")
    ]
    if not parents:
        raise ValueError("validated immediate parent required")
    parent = np.asarray(parents[-1]["solution_vector"], dtype=float)
    source = ImplicitDAEResidual(state, nodes=target_nodes - 1, uv_homotopy=1.0)
    target = ImplicitDAEResidual(state, nodes=target_nodes, uv_homotopy=1.0)
    lifts = {
        degree: endpoint_bspline_lift(source, parent, target, degree=degree)
        for degree in range(1, 6)
    }
    trials: list[dict] = []

    def evaluate(label: str, vector: np.ndarray, **parameters: float) -> None:
        projected, projection = balanced_algebraic_projection(target, vector)
        norm = float(np.max(np.abs(target.evaluate(projected))))
        trials.append({
            "label": label,
            "parameters": parameters,
            "dae_max_abs_scaled_residual": norm,
            "projection_alpha": projection["selected_alpha"],
        })

    for degree, vector in lifts.items():
        evaluate(f"bspline_degree_{degree}", vector, degree=float(degree))

    # A deliberately small bounded family.  These are diagnostic trials only;
    # any useful interval must later be encoded and regression-tested before a
    # vector can enter the persistent PASS chain.
    for lower in range(1, 5):
        for upper in range(lower + 1, 6):
            for beta in np.linspace(-3.0, 5.0, 17):
                vector = lifts[upper] + beta * (lifts[upper] - lifts[lower])
                evaluate(
                    f"degree_{lower}_{upper}_extrapolation",
                    vector,
                    lower=float(lower), upper=float(upper), beta=float(beta),
                )
    trials.sort(key=lambda item: item["dae_max_abs_scaled_residual"])
    return {
        "target_nodes": target_nodes,
        "source": f"validated_{target_nodes - 1}_PASS_only",
        "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()
