from __future__ import annotations

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_convergence import (
    JunctionPilotConfig,
    prefix_indices,
    prefix_trajectory,
    select_earliest_positive_frequency_root,
)


APEIRON_ROOT = Path(os.environ["APEIRON_AP1_ROOT"])


class AP1M1JunctionConvergenceTests(unittest.TestCase):
    def test_default_config_reserves_held_out_levels(self):
        frozen = decode_hard_pass_state(
            APEIRON_ROOT / "HANDOFF/CURRENT_STATE_LATEST.npz"
        )
        config = JunctionPilotConfig()
        config.validate(len(frozen))
        self.assertNotIn(config.held_out_momentum_nodes, config.momentum_nodes)
        self.assertEqual(config.curvature_tail_strides, (4, 2, 1))

    def test_nondivisible_prefix_is_rejected(self):
        with self.assertRaises(ValueError):
            JunctionPilotConfig(prefix_intervals=255).validate(15361)

    def test_non_nested_strides_are_rejected(self):
        with self.assertRaises(ValueError):
            JunctionPilotConfig(temporal_strides=(16, 6, 2)).validate(15361)

    def test_held_out_momentum_must_be_finer(self):
        with self.assertRaises(ValueError):
            JunctionPilotConfig(held_out_momentum_nodes=64).validate(15361)

    def test_prefix_indices_are_exactly_nested(self):
        coarse = prefix_indices(256, 16)
        fine = prefix_indices(256, 8)
        self.assertTrue(np.array_equal(fine[::2], coarse))

    def test_prefix_trajectory_uses_only_requested_HARD_PASS_prefix(self):
        frozen = decode_hard_pass_state(
            APEIRON_ROOT / "HANDOFF/CURRENT_STATE_LATEST.npz"
        )
        rows, trajectory = prefix_trajectory(frozen, 256, 16)
        self.assertEqual(len(rows), 17)
        self.assertEqual(trajectory.N[0], frozen[0, 5])
        self.assertEqual(trajectory.N[-1], frozen[256, 5])
        self.assertTrue(np.all(np.isfinite(trajectory.Hdot)))

    def test_selection_uses_earliest_admissible_root(self):
        roots = [
            {"N_local": 0.3, "positive_vacuum_frequency_margin": False},
            {"N_local": 0.2, "positive_vacuum_frequency_margin": True},
            {"N_local": 0.1, "positive_vacuum_frequency_margin": True},
        ]
        selected = select_earliest_positive_frequency_root(roots)
        self.assertEqual(selected["N_local"], 0.1)

    def test_selection_rejects_absent_admissible_root(self):
        roots = [{"N_local": 0.1, "positive_vacuum_frequency_margin": False}]
        self.assertIsNone(select_earliest_positive_frequency_root(roots))


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