"""Tests for engine.adapter — InMemoryAdapter and StorageAdapter protocol.""" import threading from engine.adapter import InMemoryAdapter from engine.types.operation import Operation # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def make_op(**kwargs: object) -> Operation: defaults: dict[str, object] = { "type": "insert", "device_id": "dev-1", "timestamp": "1000000.0", "entity_id": "e1", "entity_type": "obj", "payload": {"key": "value"}, } defaults.update(kwargs) return Operation(**defaults) # --------------------------------------------------------------------------- # Initial state # --------------------------------------------------------------------------- class TestInMemoryAdapterInit: def test_cursor_starts_at_zero(self) -> None: adapter = InMemoryAdapter() assert adapter.cursor == 0 def test_get_all_returns_empty_on_fresh_adapter(self) -> None: adapter = InMemoryAdapter() assert adapter.get_operations_after(0) == [] # --------------------------------------------------------------------------- # append_operations # --------------------------------------------------------------------------- class TestAppendOperations: def test_returns_new_cursor(self) -> None: adapter = InMemoryAdapter() cursor = adapter.append_operations([make_op(), make_op()]) assert cursor == 2 def test_empty_append_returns_current_cursor(self) -> None: adapter = InMemoryAdapter() adapter.append_operations([make_op()]) cursor = adapter.append_operations([]) assert cursor == 1 def test_sequential_ids_are_stamped(self) -> None: adapter = InMemoryAdapter() adapter.append_operations([make_op(), make_op(), make_op()]) ids = [op.id for op in adapter._log] assert ids == ["1", "2", "3"] def test_original_op_id_is_overwritten(self) -> None: """Server always assigns its own sequence number.""" adapter = InMemoryAdapter() adapter.append_operations([make_op(id="client-assigned-id")]) assert adapter._log[0].id == "1" def test_original_op_fields_preserved(self) -> None: adapter = InMemoryAdapter() op = make_op(entity_id="entity-42", payload={"hello": "world"}) adapter.append_operations([op]) stored = adapter._log[0] assert stored.entity_id == "entity-42" assert stored.payload == {"hello": "world"} def test_cursor_advances_incrementally(self) -> None: adapter = InMemoryAdapter() c1 = adapter.append_operations([make_op()]) c2 = adapter.append_operations([make_op()]) c3 = adapter.append_operations([make_op()]) assert c1 == 1 assert c2 == 2 assert c3 == 3 # --------------------------------------------------------------------------- # get_operations_after # --------------------------------------------------------------------------- class TestGetOperationsAfter: def test_cursor_zero_returns_all(self) -> None: adapter = InMemoryAdapter() adapter.append_operations([make_op(), make_op()]) ops = adapter.get_operations_after(0) assert len(ops) == 2 def test_cursor_at_end_returns_empty(self) -> None: adapter = InMemoryAdapter() adapter.append_operations([make_op(), make_op()]) ops = adapter.get_operations_after(2) assert ops == [] def test_cursor_in_middle_returns_tail(self) -> None: adapter = InMemoryAdapter() adapter.append_operations([make_op(), make_op(), make_op()]) ops = adapter.get_operations_after(1) assert len(ops) == 2 assert ops[0].id == "2" assert ops[1].id == "3" def test_returns_copy_not_reference(self) -> None: """Mutating the returned list must not affect the internal log.""" adapter = InMemoryAdapter() adapter.append_operations([make_op()]) ops = adapter.get_operations_after(0) ops.clear() assert len(adapter._log) == 1 # --------------------------------------------------------------------------- # Thread safety # --------------------------------------------------------------------------- class TestThreadSafety: def test_concurrent_appends_produce_unique_sequential_ids(self) -> None: adapter = InMemoryAdapter() errors: list[Exception] = [] def worker() -> None: try: adapter.append_operations([make_op()]) except Exception as exc: errors.append(exc) threads = [threading.Thread(target=worker) for _ in range(50)] for t in threads: t.start() for t in threads: t.join() assert errors == [], f"Thread errors: {errors}" assert adapter.cursor == 50 ids = {op.id for op in adapter._log} assert ids == {str(i) for i in range(1, 51)}