from contextlib import contextmanager import os import sys import tempfile import unittest from collections import Counter from pathlib import Path from unittest import mock ROOT = Path(__file__).resolve().parents[1] APP_DIR = ROOT / "app" sys.path.insert(0, str(APP_DIR)) import runtime_security from keycheckers.gemini import geminiKeycheck as gemini class _FailingWriteHandle: def __init__(self, handle): self._handle = handle def __enter__(self): self._handle.__enter__() return self def __exit__(self, exc_type, exc_value, traceback): return self._handle.__exit__(exc_type, exc_value, traceback) def write(self, payload): self._handle.write(payload[:max(1, len(payload) // 2)]) raise OSError("injected destination write failure") def __getattr__(self, name): return getattr(self._handle, name) class GeminiLegacyMigrationDurabilityTests(unittest.TestCase): MOVED_ONE = "AIza" + ("a" * 35) MOVED_TWO = "AIza" + ("b" * 35) RETAINED = "AIza" + ("c" * 35) TARGET_OTHER = "AIza" + ("d" * 35) @staticmethod def status_paths(root): return { status: os.path.join(root, f"{index:02d}-{status}.txt") for index, status in enumerate(gemini.STATUS_FILES) } @staticmethod def write_private(path, payload): Path(path).write_bytes(payload) runtime_security.harden_private_file(path) def seed_migration(self, paths): alive = ( f"{self.MOVED_ONE}:[legacy]:old:RATE_LIMITED\n" f"{self.RETAINED}:[current]:paid:OK\r\n" f"{self.MOVED_ONE}:[duplicate]:old:RATE_LIMITED\n" "\n" ":RATE_LIMITED\n" f"{self.MOVED_TWO}:[legacy]:free:RATE_LIMITED\n" "plain-retained-row" ).encode("utf-8") limited = ( f"{self.TARGET_OTHER}:[current]:paid:RATE_LIMITED\n" f"{self.MOVED_ONE}:[existing]:paid:RATE_LIMITED\n" f"{self.MOVED_ONE}:[duplicate-existing]:paid:RATE_LIMITED\n" ).encode("utf-8") self.write_private(paths["VALID"], alive) self.write_private(paths["VALID_RATE_LIMITED"], limited) def assert_source_or_destination(self, paths): source = gemini.load_keys_from_file(paths["VALID"]) destination = gemini.load_keys_from_file(paths["VALID_RATE_LIMITED"]) for key in (self.MOVED_ONE, self.MOVED_TWO): self.assertIn(key, source | destination) def assert_converged(self, paths): expected_alive = ( f"{self.RETAINED}:[current]:paid:OK\r\n" "\n" ":RATE_LIMITED\n" "plain-retained-row" ).encode("utf-8") self.assertEqual(Path(paths["VALID"]).read_bytes(), expected_alive) limited_lines = list(gemini.iter_bounded_text_lines(paths["VALID_RATE_LIMITED"])) counts = Counter(gemini.key_from_line(line) for line in limited_lines) self.assertEqual(counts[self.MOVED_ONE], 1) self.assertEqual(counts[self.MOVED_TWO], 1) self.assertEqual(counts[self.TARGET_OTHER], 1) self.assertEqual( limited_lines, [ f"{self.TARGET_OTHER}:[current]:paid:RATE_LIMITED\n", f"{self.MOVED_ONE}:[existing]:paid:RATE_LIMITED\n", f"{self.MOVED_TWO}:[legacy]:free:RATE_LIMITED\n", ], ) def fault_patch(self, stage, position, paths): alive_path = os.path.normcase(os.path.abspath(paths["VALID"])) limited_path = os.path.normcase(os.path.abspath(paths["VALID_RATE_LIMITED"])) if stage == "destination_write": real_writer = gemini.private_atomic_writer @contextmanager def fail_destination_write(path, *args, **kwargs): candidate = os.path.normcase(os.path.abspath(os.fspath(path))) with real_writer(path, *args, **kwargs) as handle: yield _FailingWriteHandle(handle) if candidate == limited_path else handle return mock.patch.object(gemini, "private_atomic_writer", side_effect=fail_destination_write) if stage == "destination_fsync": real_fsync = os.fsync def fail_destination_fsync(descriptor): real_fsync(descriptor) raise OSError("injected destination fsync failure") return mock.patch.object(gemini.os, "fsync", side_effect=fail_destination_fsync) writer_module = sys.modules[gemini.private_atomic_writer.__module__] real_replace = writer_module.durable_replace failed_path = limited_path if stage == "destination_replace" else alive_path def fail_replace(source, destination): candidate = os.path.normcase(os.path.abspath(os.fspath(destination))) if candidate != failed_path: return real_replace(source, destination) if position == "after": real_replace(source, destination) raise OSError(f"injected {stage} failure") return mock.patch.object(writer_module, "durable_replace", side_effect=fail_replace) def test_faults_preserve_a_copy_and_successful_rerun_converges(self): cases = ( ("destination_write", "during"), ("destination_fsync", "after"), ("destination_replace", "before"), ("destination_replace", "after"), ("source_replace", "before"), ("source_replace", "after"), ) for stage, position in cases: with self.subTest(stage=stage, position=position), tempfile.TemporaryDirectory() as temp_dir: runtime_security.ensure_private_directory(temp_dir, reject_reparse=True) paths = self.status_paths(temp_dir) self.seed_migration(paths) with mock.patch.object(gemini, "STATUS_FILES", paths): with self.fault_patch(stage, position, paths): with self.assertRaisesRegex(OSError, "injected"): gemini.migrate_legacy_alive_rate_limited() self.assert_source_or_destination(paths) gemini.migrate_legacy_alive_rate_limited() self.assert_converged(paths) def test_no_migration_is_unchanged_and_uses_bounded_status_lock(self): with tempfile.TemporaryDirectory() as temp_dir: runtime_security.ensure_private_directory(temp_dir, reject_reparse=True) paths = self.status_paths(temp_dir) alive = f"{self.RETAINED}:[current]:paid:OK\nplain-retained-row".encode("utf-8") limited = f"{self.TARGET_OTHER}:[current]:paid:RATE_LIMITED\n".encode("utf-8") self.write_private(paths["VALID"], alive) self.write_private(paths["VALID_RATE_LIMITED"], limited) real_acquire = gemini.acquire_file_lock with mock.patch.object(gemini, "STATUS_FILES", paths), \ mock.patch.object(gemini, "acquire_file_lock", wraps=real_acquire) as acquire, \ mock.patch.object(gemini, "_replace_status_snapshot") as replace: gemini.migrate_legacy_alive_rate_limited() self.assertEqual(Path(paths["VALID"]).read_bytes(), alive) self.assertEqual(Path(paths["VALID_RATE_LIMITED"]).read_bytes(), limited) replace.assert_not_called() acquire.assert_called_once_with( os.path.join(temp_dir, "geminiStatus.lock"), timeout_sec=30, ) if __name__ == "__main__": unittest.main()