from __future__ import annotations

from dataclasses import replace
import json
from pathlib import Path
import tempfile
import unittest

import numpy as np

import ap1_m1_coupled_moving_split_background as coupled


ROOT = Path(__file__).resolve().parents[2]


class CoupledMovingSplitBackgroundTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls) -> None:
        cls.seed_report, cls.arrays, cls.authorities = coupled.load_inputs(ROOT)
        cls.p = coupled.ChiParameters()

    def short_trajectory(self, span: float = 2.5e-4, nodes: int = 33):
        grid = np.linspace(0.0, span, nodes)
        stored = np.asarray(self.arrays["quantum_rho_pressure_chi2"], dtype=float)
        source = {
            key: np.full(nodes, stored[index])
            for index, key in enumerate(coupled.SOURCE_KEYS)
        }
        return coupled.integrate_background(self.arrays, grid, source, 2.0e-9)[0]

    def test_authority_chain_is_exact_and_pass(self) -> None:
        self.assertTrue(self.seed_report["all_seed_gates_pass"])
        self.assertEqual(self.authorities, coupled.EXPECTED_AUTHORITIES)

    def test_frozen_split_and_pv_basis(self) -> None:
        config = coupled.CoupledConfig()
        config.validate()
        self.assertEqual(config.qcut_over_Lambda, 0.6)
        np.testing.assert_array_equal(coupled.PV_C, [1.0, -3.0, 3.0, -1.0])
        np.testing.assert_array_equal(coupled.PV_J, [0.0, 1.0, 2.0, 3.0])

    def test_config_rejects_changed_split(self) -> None:
        with self.assertRaises(ValueError):
            replace(coupled.CoupledConfig(), qcut_over_Lambda=0.61).validate()

    def test_local_convective_flux_matches_independent_fixed_k(self) -> None:
        trajectory = self.short_trajectory()
        qcut = 0.6 * self.p.Lambda
        index = len(trajectory.N) // 2
        local = coupled.adiabatic_local_table(
            trajectory, np.array([qcut]), 0.5 * qcut, 2.0 * qcut, 129, self.p
        )
        independent = coupled.adiabatic_primitives_at(
            trajectory, float(trajectory.N[index]), np.array([qcut]), self.p
        )
        primitive = {
            "q": np.array([qcut]),
            "H": float(trajectory.H[index]),
            **{
                name: local[name][index]
                for name in (
                    "mass2", "omega", "omega_dot", "W2", "W2_dot",
                    "W4", "W", "W_dot",
                )
            },
        }
        observed = coupled.expanded_tail_flux_longdouble(primitive)
        reference = coupled.expanded_tail_flux_longdouble(independent)
        maximum = 0.0
        for key in coupled.SOURCE_KEYS:
            left = observed[key][:, 0]
            right = reference[key][:, 0]
            delta = np.sum(coupled.PV_C.astype(np.longdouble) * (left - right))
            gross = max(
                np.sum(np.abs(coupled.PV_C * left)),
                np.sum(np.abs(coupled.PV_C * right)),
                np.longdouble(1e-300),
            )
            maximum = max(maximum, float(abs(delta) / gross))
        self.assertLess(maximum, coupled.FREEZE_POLICY["local_fixed_k_source_gross_relative_cap"])

    def test_handoff_table_preserves_canonical_state(self) -> None:
        trajectory = self.short_trajectory()
        qcut = 0.6 * self.p.Lambda
        table = coupled.boundary_handoff_table(
            trajectory, len(trajectory.N), qcut, self.p, 129
        )
        _, audit = coupled._handoff_state_from_table(
            table, float(trajectory.N[len(trajectory.N) // 2])
        )
        self.assertGreater(audit["min_omega2"], 0.0)
        self.assertGreater(audit["min_W"], 0.0)
        self.assertLess(audit["state_transfer_relative_error"], 1.0e-12)
        self.assertLess(audit["initial_wronskian_relative_error"], 1.0e-12)

    def test_moving_tail_is_finite_and_boundary_closed(self) -> None:
        trajectory = self.short_trajectory()
        tail = coupled.moving_adiabatic_tail(
            trajectory, 2, 8.0,
            adiabatic_support_nodes=len(trajectory.N),
            adiabatic_momentum_support_nodes=129,
        )
        self.assertGreater(tail["diagnostics"]["min_tail_omega2_Mpl2"], 0.0)
        self.assertGreater(tail["diagnostics"]["min_tail_W_Mpl"], 0.0)
        self.assertLess(
            tail["diagnostics"]["max_boundary_direct_expanded_gross_relative_delta"],
            coupled.FREEZE_POLICY["boundary_direct_expanded_gross_relative_cap"],
        )
        for key in coupled.SOURCE_KEYS:
            self.assertTrue(np.all(np.isfinite(tail["signed"][key])))

    def test_resolved_crossing_count_and_wronskian(self) -> None:
        trajectory = self.short_trajectory(span=1.0e-4, nodes=9)
        result = coupled.transport_resolved_modes(
            trajectory, self.arrays, 1, 2.5e-5,
            adiabatic_support_nodes=9,
            adiabatic_momentum_support_nodes=65,
        )
        diagnostics = result["diagnostics"]
        self.assertEqual(diagnostics["new_entry_quadrature_modes"], 8)
        self.assertEqual(diagnostics["final_resolved_modes"], 136)
        self.assertEqual(diagnostics["vacuum_resets"], 0)
        self.assertLess(diagnostics["max_wronskian_relative_error"], 1.0e-11)

    def test_background_starts_from_exact_seed(self) -> None:
        trajectory = self.short_trajectory(span=1.0e-4, nodes=9)
        state = np.asarray(self.arrays["m1_state_physical_N"], dtype=float)
        self.assertEqual(trajectory.sigma[0], state[0])
        self.assertEqual(trajectory.theta[0], state[2])
        self.assertEqual(trajectory.H[0], state[4])
        self.assertEqual(trajectory.Hdot[0], float(self.arrays["Hdot_Mpl2"][0]))

    def test_bounded_coupled_smoke_has_no_reset_or_rows(self) -> None:
        config = replace(coupled.CoupledConfig(), fixed_point_steps=2)
        run = coupled.run_coupled(
            self.arrays, 1.0e-4, 9, config,
            background_rtol=2.0e-9,
            max_mode_step=2.5e-5,
            entry_nodes_per_panel=1,
            tail_nodes_per_octave=2,
            K_over_Lambda=8.0,
            adiabatic_support_nodes=9,
            adiabatic_momentum_support_nodes=65,
        )
        self.assertTrue(np.all(run["trajectory"].H > 0.0))
        self.assertEqual(run["source_diagnostics"]["resolved"]["vacuum_resets"], 0)
        self.assertTrue(run["source_diagnostics"]["junction_values_enforced_exactly"])

    def test_nonpass_report_is_never_written(self) -> None:
        with tempfile.TemporaryDirectory() as directory:
            target = Path(directory) / "forbidden.json"
            with self.assertRaises(coupled.CoupledBackgroundError):
                coupled._write_pass_json({"all_pass": False}, target, "all_pass")
            self.assertFalse(target.exists())

    def test_report_payload_is_json_serializable(self) -> None:
        payload = {
            "policy": coupled.FREEZE_POLICY,
            "config": coupled.asdict(coupled.CoupledConfig()),
        }
        self.assertIsInstance(json.dumps(payload), str)


if __name__ == "__main__":
    unittest.main()
