#!/usr/bin/env python3 """Unit tests for braid_diat_codec.py — encode/decode roundtrip and layer logic.""" import importlib.util import struct import sys import unittest from pathlib import Path MODULE_PATH = Path(__file__).with_name("braid_diat_codec.py") SPEC = importlib.util.spec_from_file_location("braid_diat_codec", MODULE_PATH) assert SPEC is not None and SPEC.loader is not None codec = importlib.util.module_from_spec(SPEC) sys.modules[SPEC.name] = codec SPEC.loader.exec_module(codec) # --------------------------------------------------------------------------- # Q0_2 encoding # --------------------------------------------------------------------------- class TestQ02Encoding(unittest.TestCase): def test_encode_all_states(self): self.assertEqual(codec.encode_q02(0), 0) self.assertEqual(codec.encode_q02(16384), 1) self.assertEqual(codec.encode_q02(32768), 2) self.assertEqual(codec.encode_q02(49152), 3) def test_decode_all_states(self): self.assertEqual(codec.decode_q02(0), 0) self.assertEqual(codec.decode_q02(1), 16384) self.assertEqual(codec.decode_q02(2), 32768) self.assertEqual(codec.decode_q02(3), 49152) def test_roundtrip_all_states(self): for v in codec.Q0_2_STATES: self.assertEqual(codec.decode_q02(codec.encode_q02(v)), v) def test_decode_masks_to_2_bits(self): self.assertEqual(codec.decode_q02(4), 0) self.assertEqual(codec.decode_q02(5), 16384) # --------------------------------------------------------------------------- # Layer 1: ChiralityDIAT # --------------------------------------------------------------------------- class TestChiralityDIAT(unittest.TestCase): def test_encode_decode_roundtrip(self): original = codec.ChiralityDIAT(codec.Chirality.POSITIVE, 5, 100, 200, 42) raw = original.to_bytes() self.assertEqual(len(raw), 8) restored = codec.ChiralityDIAT.from_bytes(raw) self.assertEqual(restored.chirality, codec.Chirality.POSITIVE) self.assertEqual(restored.shell, 5) self.assertEqual(restored.offset_a, 100) self.assertEqual(restored.offset_b, 200) self.assertEqual(restored.prod_msb, 42) def test_all_chiralities(self): for ch in codec.Chirality: obj = codec.ChiralityDIAT(ch, 0, 0, 0, 0) restored = codec.ChiralityDIAT.from_bytes(obj.to_bytes()) self.assertEqual(restored.chirality, ch) def test_decode_from_n(self): slot = codec.ChiralityDIAT.decode(codec.Chirality.LEFT, 25) self.assertIsNotNone(slot) self.assertEqual(slot.chirality, codec.Chirality.LEFT) self.assertEqual(slot.shell, 5) self.assertEqual(slot.offset_a, 0) self.assertEqual(slot.offset_b, 11) def test_decode_n_too_large_returns_none(self): result = codec.ChiralityDIAT.decode(codec.Chirality.NONE, 0x400000) self.assertIsNone(result) def test_to_n_roundtrip(self): for n in [0, 1, 4, 9, 16, 25, 100, 1000]: slot = codec.ChiralityDIAT.decode(codec.Chirality.POSITIVE, n) self.assertIsNotNone(slot) self.assertEqual(slot.to_n(), n) def test_verify_b(self): slot = codec.ChiralityDIAT.decode(codec.Chirality.ACHIRAL, 17) self.assertIsNotNone(slot) self.assertTrue(slot.verify_b()) def test_verify_b_corrupted(self): slot = codec.ChiralityDIAT.decode(codec.Chirality.ACHIRAL, 17) slot.offset_b = 999 self.assertFalse(slot.verify_b()) def test_bytes_roundtrip_preserves_all_fields(self): for n in [0, 7, 50, 255]: slot = codec.ChiralityDIAT.decode(codec.Chirality.POSITIVE, n) restored = codec.ChiralityDIAT.from_bytes(slot.to_bytes()) self.assertEqual(slot.shell, restored.shell) self.assertEqual(slot.offset_a, restored.offset_a) self.assertEqual(slot.offset_b, restored.offset_b) self.assertEqual(slot.prod_msb, restored.prod_msb) # --------------------------------------------------------------------------- # Layer 2: MountainPacked # --------------------------------------------------------------------------- class TestMountainPacked(unittest.TestCase): def _sample_mountain_dict(self): return { "height": 3, "apex": [100, 200, 300], "base": [[10, 20, 30], [40, 50, 60]], "inner": {}, } def test_from_mountain_roundtrip(self): m = self._sample_mountain_dict() packed = codec.MountainPacked.from_mountain(m) self.assertEqual(packed.height, 3) self.assertEqual(packed.apex_x, 100) self.assertEqual(packed.apex_y, 200) self.assertEqual(packed.apex_z, 300) self.assertEqual(packed.base_count, 2) result = packed.to_mountain() self.assertEqual(result["height"], 3) self.assertEqual(result["apex"], [100, 200, 300]) self.assertEqual(result["base"], [[10, 20, 30], [40, 50, 60]]) def test_bytes_roundtrip(self): m = self._sample_mountain_dict() packed = codec.MountainPacked.from_mountain(m) raw = packed.to_bytes() restored = codec.MountainPacked.from_bytes(raw) self.assertEqual(restored.height, packed.height) self.assertEqual(restored.apex_x, packed.apex_x) self.assertEqual(restored.apex_y, packed.apex_y) self.assertEqual(restored.apex_z, packed.apex_z) self.assertEqual(restored.base_count, packed.base_count) self.assertEqual(restored.bases, packed.bases) def test_empty_bases(self): m = {"height": 1, "apex": [0, 0, 0], "base": [], "inner": {}} packed = codec.MountainPacked.from_mountain(m) self.assertEqual(packed.base_count, 0) self.assertEqual(packed.bases, []) result = packed.to_mountain() self.assertEqual(result["base"], []) def test_negative_coords(self): m = {"height": 2, "apex": [-100, -200, -300], "base": [[-1, -2, -3]], "inner": {}} packed = codec.MountainPacked.from_mountain(m) raw = packed.to_bytes() restored = codec.MountainPacked.from_bytes(raw) result = restored.to_mountain() self.assertEqual(result["apex"], [-100, -200, -300]) self.assertEqual(result["base"], [[-1, -2, -3]]) # --------------------------------------------------------------------------- # Layer 3: BraidResidualPacked # --------------------------------------------------------------------------- class TestBraidResidualPacked(unittest.TestCase): def _sample_bracket(self): return { "lower": 0, "upper": 16384, "gap": 32768, "kappa": 49152, "phi": 0, "admissible": True, } def test_from_bracket_roundtrip(self): br = self._sample_bracket() packed = codec.BraidResidualPacked.from_bracket(br) result = packed.to_bracket() self.assertEqual(result, br) def test_bytes_roundtrip(self): br = self._sample_bracket() packed = codec.BraidResidualPacked.from_bracket(br) raw = packed.to_bytes() self.assertEqual(len(raw), 8) restored = codec.BraidResidualPacked.from_bytes(raw) self.assertEqual(restored.to_bracket(), br) def test_all_q02_states_roundtrip(self): for v in codec.Q0_2_STATES: br = {"lower": v, "upper": v, "gap": v, "kappa": v, "phi": v, "admissible": False} packed = codec.BraidResidualPacked.from_bracket(br) result = codec.BraidResidualPacked.from_bytes(packed.to_bytes()).to_bracket() self.assertEqual(result, br) def test_admissible_bit(self): br_true = self._sample_bracket() br_true["admissible"] = True br_false = dict(br_true) br_false["admissible"] = False packed_t = codec.BraidResidualPacked.from_bracket(br_true) packed_f = codec.BraidResidualPacked.from_bracket(br_false) self.assertTrue(codec.BraidResidualPacked.from_bytes(packed_t.to_bytes()).admissible) self.assertFalse(codec.BraidResidualPacked.from_bytes(packed_f.to_bytes()).admissible) # --------------------------------------------------------------------------- # Chirality enum # --------------------------------------------------------------------------- class TestChiralityEnum(unittest.TestCase): def test_values(self): self.assertEqual(int(codec.Chirality.NONE), 0) self.assertEqual(int(codec.Chirality.POSITIVE), 1) self.assertEqual(int(codec.Chirality.LEFT), 2) self.assertEqual(int(codec.Chirality.ACHIRAL), 3) def test_all_fit_in_2_bits(self): for ch in codec.Chirality: self.assertLessEqual(int(ch), 3) if __name__ == "__main__": unittest.main()