from __future__ import annotations

from copy import deepcopy
import json
import os
from pathlib import Path
import unittest

import numpy as np

from ap1_m1_background_preflight import decode_hard_pass_state
from ap1_m1_junction_seed import (
    JunctionSeedError,
    build_seed_payload,
    canonical_mode_state,
    validate_seed_authorities,
)


APEIRON_ROOT = Path(os.environ["APEIRON_AP1_ROOT"])
HELDOUT_PATH = APEIRON_ROOT / "AP1/APEIRON_AP1_M1_EXACT_INITIAL_JUNCTION_HELDOUT_LATEST.json"
MANIFEST_PATH = APEIRON_ROOT / "AP1/APEIRON_AP1_M1_EXACT_INITIAL_JUNCTION_TOLERANCES_LATEST.json"
PILOT_PATH = APEIRON_ROOT / "AP1/APEIRON_AP1_M1_JUNCTION_CONVERGENCE_PILOT_LATEST.json"
CONVERGENCE_PATH = APEIRON_ROOT / "AP1/CODE/ap1_m1_junction_convergence.py"
BOUNDARY_PATH = APEIRON_ROOT / "AP1/CODE/ap1_m1_boundary_match.py"
HELDOUT_CODE_PATH = APEIRON_ROOT / "AP1/CODE/ap1_m1_junction_heldout.py"


class AP1M1JunctionSeedTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.heldout = json.loads(HELDOUT_PATH.read_text(encoding="utf-8"))
        cls.frozen = decode_hard_pass_state(
            APEIRON_ROOT / "HANDOFF/CURRENT_STATE_LATEST.npz"
        )
        cls.mode = canonical_mode_state(cls.frozen[0], 96)

    def test_heldout_authorities_validate(self):
        observed = validate_seed_authorities(
            APEIRON_ROOT,
            self.heldout,
            HELDOUT_PATH,
            MANIFEST_PATH,
            PILOT_PATH,
            CONVERGENCE_PATH,
            BOUNDARY_PATH,
            HELDOUT_CODE_PATH,
        )
        self.assertGreaterEqual(len(observed), 13)

    def test_nonpass_heldout_is_rejected(self):
        bad = deepcopy(self.heldout)
        bad["all_exact_initial_junction_gates_pass"] = False
        with self.assertRaises(JunctionSeedError):
            validate_seed_authorities(
                APEIRON_ROOT,
                bad,
                HELDOUT_PATH,
                MANIFEST_PATH,
                PILOT_PATH,
                CONVERGENCE_PATH,
                BOUNDARY_PATH,
                HELDOUT_CODE_PATH,
            )

    def test_canonical_mode_shapes_are_4_by_96(self):
        for key in ("omega2_Mpl2", "u_real", "u_imag", "v_real", "v_imag"):
            self.assertEqual(self.mode[key].shape, (4, 96))
            self.assertTrue(np.all(np.isfinite(self.mode[key])))

    def test_all_mode_frequencies_are_positive(self):
        self.assertGreater(float(np.min(self.mode["omega2_Mpl2"])), 0.0)

    def test_canonical_initial_wronskian(self):
        error = np.max(
            np.abs(self.mode["wronskian_observed"] - self.mode["wronskian_target"])
        )
        self.assertLess(error, 1.0e-14)

    def test_invalid_state_is_rejected(self):
        bad = self.frozen[0].copy()
        bad[0] = np.nan
        with self.assertRaises(ValueError):
            canonical_mode_state(bad, 96)

    def test_seed_report_is_json_serializable(self):
        arrays, report = build_seed_payload(APEIRON_ROOT, self.heldout, {})
        self.assertTrue(report["all_seed_gates_pass"])
        self.assertEqual(len(arrays["k_comoving_Mpl"]), 96)
        json.dumps(report)


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