"""Implicit Lobatto DAE residual assembly for the AP1-R2c candidate."""
from __future__ import annotations

from dataclasses import asdict
import json
from pathlib import Path

import numpy as np

from ap1_m1_background_preflight import decode_hard_pass_state
from ap1_r2c_high_precision_modes import completed_history_decimal, propagate_endpoint_decimal
from ap1_r2c_implicit_dae_preflight import lobatto_grid_and_D, positive_friedmann_root
from ap1_r2c_mode_inheritance import inherited_trajectory
from ap1_r2c_self_consistent_candidate import (
    CandidateConfig, classical_terms, integrate_segment, match_seed_N, standard_stress,
)
from chi_background_closure import (
    ChiParameters, ChiTrajectory, PV_C, PV_J, _logarithmic_shells,
    frozen_renormalization_constants, order0_uv_tail_n, physical_shells, portal_terms,
)
from planck2018_neutrino_closure import Planck2018Pilot


X_KEYS = ("sigma", "sigma_dot", "theta", "theta_dot")
Q_KEYS = ("rho", "pressure", "chi2")


def _Dt(values: np.ndarray, H: np.ndarray, D_N: np.ndarray) -> np.ndarray:
    differentiated = D_N @ np.asarray(values)
    shape = (len(H),) + (1,) * (differentiated.ndim - 1)
    return H.reshape(shape) * differentiated


def curvature_uv_tail_collocation(
    trajectory: ChiTrajectory, D_N: np.ndarray,
    p: ChiParameters | None = None, nodes_per_octave: int = 6,
) -> dict[str, np.ndarray]:
    """Canonical adiabatic orders 2+4 using one global DAE derivative."""
    p = ChiParameters() if p is None else p
    tr = trajectory.validated()
    if D_N.shape != (len(tr.N), len(tr.N)):
        raise ValueError("DAE derivative matrix must match trajectory")
    a = np.exp(tr.N)
    kcut = 0.6 * p.Lambda
    K = 64.0 * p.Lambda * float(np.max(a))
    k, weights, _segments = _logarithmic_shells(kcut, K, nodes_per_octave)
    physical_k2 = k[None, :] ** 2 / a[:, None] ** 2
    mass2 = np.asarray(portal_terms(tr.sigma, tr.theta, p)[0])
    muR2 = (p.muR_over_Lambda * p.Lambda) ** 2
    rho = np.zeros(len(tr.N)); pressure = np.zeros(len(tr.N)); chi2 = np.zeros(len(tr.N))
    for sector in range(4):
        masses2 = mass2[:, None] + PV_J[sector] * muR2
        omega2 = physical_k2 + masses2
        if np.min(omega2) <= 0.0:
            raise ValueError("collocation curvature tail entered omega2 <= 0")
        omega = np.sqrt(omega2)
        omega_dot = _Dt(omega, tr.H, D_N)
        omega_ddot = _Dt(omega_dot, tr.H, D_N)
        logarithmic0 = omega_dot / omega
        s2 = -2.25 * tr.H[:, None] ** 2 - 1.5 * tr.Hdot[:, None]
        q2 = 0.75 * logarithmic0**2 - 0.5 * omega_ddot / omega
        W2 = (s2 + q2) / (2.0 * omega)
        W2_dot = _Dt(W2, tr.H, D_N)
        W2_ddot = _Dt(W2_dot, tr.H, D_N)
        logarithmic2 = W2_dot / omega - logarithmic0 * W2 / omega
        q4 = 1.5 * logarithmic0 * logarithmic2 - 0.5 * (
            W2_ddot / omega - omega_ddot * W2 / omega**2
        )
        W4 = (q4 - W2**2) / (2.0 * omega)
        prefactor = 1.0 / (2.0 * a[:, None] ** 3)
        A0 = prefactor / omega
        A2 = -prefactor * W2 / omega**2
        A4 = prefactor * (W2**2 / omega**3 - W4 / omega**2)
        B1 = 1.5 * tr.H[:, None] + 0.5 * logarithmic0
        B3 = 0.5 * logarithmic2
        C0 = omega2
        C2 = 2.0 * omega * W2 + B1**2
        C4 = W2**2 + 2.0 * omega * W4 + 2.0 * B1 * B3
        D2 = A0 * C2 + A2 * C0
        D4 = A0 * C4 + A2 * C2 + A4 * C0
        rho24 = 0.5 * (D2 + omega2 * A2 + D4 + omega2 * A4)
        pressure_weight = physical_k2 / 6.0 + 0.5 * masses2
        pressure24 = 0.5 * (D2 + D4) - pressure_weight * (A2 + A4)
        coefficient = PV_C[sector]
        rho += coefficient * np.sum(weights[None, :] * rho24, axis=1)
        pressure += coefficient * np.sum(weights[None, :] * pressure24, axis=1)
        chi2 += coefficient * np.sum(weights[None, :] * (A2 + A4), axis=1)
    return {"rho": rho, "pressure": pressure, "chi2": chi2}


class ImplicitDAEResidual:
    def __init__(self, state_path: Path, nodes: int = 17, delta_N: float = 1.0e-7):
        self.cfg = CandidateConfig(delta_N=delta_N, candidate_nodes=nodes,
                                   inherited_nodes=1025, fixed_point_steps=1)
        self.Nrel, self.D_N = lobatto_grid_and_D(nodes, delta_N)
        self.pilot = Planck2018Pilot(); self.chi_p = ChiParameters()
        frozen = decode_hard_pass_state(state_path)
        self.seed = frozen[-1, :5].copy()
        self.prefix = inherited_trajectory(state_path, self.cfg.inherited_nodes)
        self.k, self.weights = physical_shells(self.chi_p.Lambda, nodes=self.cfg.mode_nodes)
        inherited = completed_history_decimal(
            self.prefix, self.k, self.weights, self.cfg.digits
        )["completed_history"]
        self.q0 = {key: np.full(nodes, float(inherited[key][-1])) for key in Q_KEYS}
        self.N_seed0 = match_seed_N(self.seed, float(self.q0["rho"][0]), self.pilot)
        y, _Hdot = integrate_segment(
            self.seed, self.Nrel, self.N_seed0, self.q0, self.pilot, self.cfg
        )
        self.X0 = {key: y[:, i].copy() for i, key in enumerate(X_KEYS)}
        self.x_scales = {
            key: max(float(np.max(np.abs(value))), 1.0e-12)
            for key, value in self.X0.items()
        }
        self.q_scales = {
            key: max(float(np.max(np.abs(value))), 1.0e-300)
            for key, value in self.q0.items()
        }
        self.N_seed_scale = 100.0
        self.initial = self.pack(self.X0, self.q0, self.N_seed0)
        self.last_blocks: dict[str, float] = {}

    def pack(self, X: dict[str, np.ndarray], Q: dict[str, np.ndarray], N_seed: float) -> np.ndarray:
        parts = [np.asarray(X[key]) / self.x_scales[key] for key in X_KEYS]
        parts += [np.asarray(Q[key]) / self.q_scales[key] for key in Q_KEYS]
        parts.append(np.array([N_seed / self.N_seed_scale]))
        return np.concatenate(parts)

    def unpack(self, vector: np.ndarray) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray], float]:
        n = self.cfg.candidate_nodes; vector = np.asarray(vector)
        X = {key: vector[i*n:(i+1)*n] * self.x_scales[key]
             for i, key in enumerate(X_KEYS)}
        offset = len(X_KEYS) * n
        Q = {key: vector[offset+i*n:offset+(i+1)*n] * self.q_scales[key]
             for i, key in enumerate(Q_KEYS)}
        return X, Q, float(vector[-1] * self.N_seed_scale)

    def _sources(self, X: dict[str, np.ndarray], H: np.ndarray, Hdot: np.ndarray) -> tuple[dict[str, np.ndarray], float]:
        candidate = ChiTrajectory(
            N=self.prefix.N[-1] + self.Nrel, H=H, Hdot=Hdot,
            sigma=X["sigma"], theta=X["theta"],
        )
        full = ChiTrajectory(
            N=np.concatenate([self.prefix.N, candidate.N[1:]]),
            H=np.concatenate([self.prefix.H, H[1:]]),
            Hdot=np.concatenate([self.prefix.Hdot, Hdot[1:]]),
            sigma=np.concatenate([self.prefix.sigma, X["sigma"][1:]]),
            theta=np.concatenate([self.prefix.theta, X["theta"][1:]]),
        )
        modes = propagate_endpoint_decimal(
            full, self.k, self.weights, self.cfg.digits, record_history=True
        )
        n = self.cfg.candidate_nodes
        resolved = {key: np.asarray(modes["resolved_pv_history"][key][-n:]) for key in Q_KEYS}
        tail0 = order0_uv_tail_n(candidate, self.chi_p)
        tail24 = curvature_uv_tail_collocation(candidate, self.D_N, self.chi_p)
        constants = frozen_renormalization_constants(self.chi_p)
        mass2 = np.asarray(portal_terms(X["sigma"], X["theta"], self.chi_p)[0])
        C=constants["C_chi2"]; A=constants["A_g"]; B=constants["B_G_rhs"]
        fresh = {
            "rho": resolved["rho"]+tail0["rho"]+tail24["rho"]-0.5*C*mass2+A+3.0*B*H**2,
            "pressure": resolved["pressure"]+tail0["pressure"]+tail24["pressure"]+0.5*C*mass2-A-B*(2.0*Hdot+3.0*H**2),
            "chi2": resolved["chi2"]+tail0["chi2"]+tail24["chi2"]-C,
        }
        return fresh, float(modes["wronskian_relative_error"])

    def evaluate(self, vector: np.ndarray) -> np.ndarray:
        X, Q, N_seed = self.unpack(vector)
        states_without_H = np.column_stack([X[key] for key in X_KEYS])
        rho_cl = np.array([classical_terms(np.r_[row, self.seed[4]])["rho"]
                           for row in states_without_H])
        rho_std, _p_std = standard_stress(N_seed + self.Nrel, self.pilot)
        H = positive_friedmann_root(rho_cl + rho_std + Q["rho"])
        Hdot = _Dt(H, H, self.D_N)
        fresh, wronskian = self._sources(X, H, Hdot)

        states = np.column_stack([states_without_H, H])
        dlist = [classical_terms(row) for row in states]
        if min(d["PX"] for d in dlist) <= 0.0 or min(d["K"] for d in dlist) <= 0.0:
            raise ValueError("DAE trial left positive local kinetic branch")
        mass = [portal_terms(row[0], row[2], self.chi_p) for row in states]
        rhs = {key: np.empty(self.cfg.candidate_nodes) for key in X_KEYS}
        rhs["sigma"] = X["sigma_dot"] / H
        rhs["theta"] = X["theta_dot"] / H
        rhs["sigma_dot"] = np.array([
            (-3.0*H[i]*X["sigma_dot"][i] + dlist[i]["P_sigma"]
             - 0.5*float(mass[i][1])*Q["chi2"][i]) / H[i]
            for i in range(self.cfg.candidate_nodes)
        ])
        rhs["theta_dot"] = np.array([
            (dlist[i]["P_theta"] - 0.5*float(mass[i][2])*Q["chi2"][i]
             - 3.0*H[i]*dlist[i]["PX"]*X["theta_dot"][i]
             - dlist[i]["PX_sigma"]*X["sigma_dot"][i]*X["theta_dot"][i])
            / (dlist[i]["K"]*H[i])
            for i in range(self.cfg.candidate_nodes)
        ])
        x_residuals = []
        seed_values = dict(zip(X_KEYS, self.seed[:4]))
        for key in X_KEYS:
            residual = self.cfg.delta_N * (self.D_N @ X[key] - rhs[key]) / self.x_scales[key]
            residual[0] = (X[key][0] - seed_values[key]) / self.x_scales[key]
            x_residuals.append(residual)
        q_residuals = [(Q[key] - fresh[key]) / self.q_scales[key] for key in Q_KEYS]
        junction = (3.0*self.seed[4]**2 - classical_terms(self.seed)["rho"]
                    - float(standard_stress(np.array([N_seed]), self.pilot)[0][0])
                    - Q["rho"][0]) / (3.0*self.seed[4]**2)
        result = np.concatenate([*x_residuals, *q_residuals, np.array([junction])])
        self.last_blocks = {
            "background": float(max(np.max(np.abs(r)) for r in x_residuals)),
            "source": float(max(np.max(np.abs(r)) for r in q_residuals)),
            "junction": float(abs(junction)), "wronskian": wronskian,
            "min_H": float(np.min(H)), "min_PX": float(min(d["PX"] for d in dlist)),
            "min_K": float(min(d["K"] for d in dlist)),
        }
        return result


def build_residual_report(state_path: Path) -> dict:
    system = ImplicitDAEResidual(state_path)
    residual = system.evaluate(system.initial)
    gates = {
        "dimension_120": len(residual) == 120,
        "finite_initial_residual": bool(np.all(np.isfinite(residual))),
        "positive_initial_H": system.last_blocks["min_H"] > 0.0,
        "positive_initial_PX_and_K": system.last_blocks["min_PX"] > 0.0 and system.last_blocks["min_K"] > 0.0,
        "wronskian_below_1e_11": system.last_blocks["wronskian"] < 1.0e-11,
    }
    return {
        "schema": "apeiron-ap1-r2c-implicit-dae-residual-v1.0",
        "classification": "FULL_DAE_RESIDUAL_ASSEMBLED_SOLVE_NOT_RUN" if all(gates.values()) else "DAE_RESIDUAL_ASSEMBLY_NOT_PASS",
        "config": asdict(system.cfg), "initial_residual_blocks": system.last_blocks,
        "initial_max_abs_scaled_residual": float(np.max(np.abs(residual))),
        "gates": gates, "old_solver_or_physical_map_called": False,
        "next_required": "safeguarded Newton-Krylov solve on 17 nodes",
        "claim_boundary": "residual assembly only; no DAE solution, candidate PASS, or observable released",
    }


def main() -> None:
    import argparse
    parser=argparse.ArgumentParser(); parser.add_argument("state",type=Path); parser.add_argument("--output",type=Path)
    args=parser.parse_args(); rendered=json.dumps(build_residual_report(args.state),indent=2)+"\n"
    if args.output: args.output.write_text(rendered,encoding="utf-8")
    else: print(rendered,end="")


if __name__ == "__main__": main()
