from __future__ import annotations

import unittest

from ap1_m1_junction_heldout import (
    METRIC_KEYS,
    REPRESENTATION_FLOORS,
    JunctionHeldoutError,
    delta_gates,
    heldout_caps,
)


class AP1M1JunctionHeldoutTests(unittest.TestCase):
    def test_zero_pilot_deltas_receive_independent_floors(self):
        caps = heldout_caps({key: 0.0 for key in METRIC_KEYS})
        self.assertEqual(caps, REPRESENTATION_FLOORS)

    def test_nonzero_pilot_delta_is_not_relaxed(self):
        values = {key: 0.0 for key in METRIC_KEYS}
        values["abs_delta_N_physical"] = 2.0e-6
        caps = heldout_caps(values)
        self.assertEqual(caps["abs_delta_N_physical"], 2.0e-6)

    def test_missing_metric_fails_closed(self):
        with self.assertRaises(JunctionHeldoutError):
            heldout_caps({"abs_delta_N_physical": 0.0})

    def test_delta_gate_accepts_cap_and_rejects_excess(self):
        caps = heldout_caps({key: 0.0 for key in METRIC_KEYS})
        observed = dict(caps)
        self.assertTrue(all(delta_gates(observed, caps).values()))
        observed["abs_delta_N_physical"] *= 2.0
        self.assertFalse(
            delta_gates(observed, caps)[
                "heldout_abs_delta_N_physical_within_cap"
            ]
        )


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