"""Boundary contracts for deterministic and transactional checkpoint state.""" from __future__ import annotations import json import os from concurrent.futures import ThreadPoolExecutor from pathlib import Path from threading import Barrier from unittest.mock import MagicMock import pytest import torch import obliteratus.persistence_contracts as persistence pytestmark = pytest.mark.cpu def _write_valid_local_checkpoint( checkpoint_dir: Path, metadata_json: str = '{"schema": 1}', ) -> None: (checkpoint_dir / "abliteration_metadata.json").write_text( metadata_json, encoding="utf-8", ) (checkpoint_dir / "config.json").write_text("{}", encoding="utf-8") (checkpoint_dir / "tokenizer_config.json").write_text("{}", encoding="utf-8") (checkpoint_dir / "model.safetensors").write_bytes(b"weights") @pytest.mark.parametrize( ("state_dict", "expected"), [ ({}, 0), ({"float": torch.ones(3, dtype=torch.float32)}, 12), ( { "half": torch.ones(5, dtype=torch.float16), "index": torch.ones(2, dtype=torch.int64), }, 26, ), ], ) def test_state_dict_size_bytes_is_exact(state_dict, expected): assert persistence.state_dict_size_bytes(state_dict) == expected @pytest.mark.parametrize( ("payload", "required"), [(0, 0), (1, 2), (9, 10), (10, 11), (101, 112)], ) def test_required_checkpoint_bytes_rounds_headroom_up(payload, required): assert persistence.required_checkpoint_bytes(payload) == required def test_required_checkpoint_bytes_rejects_negative_payload(): with pytest.raises(ValueError) as exc_info: persistence.required_checkpoint_bytes(-1) assert str(exc_info.value) == "checkpoint payload size cannot be negative" def test_checkpoint_capacity_accepts_exact_boundary_and_rejects_one_byte_less(): persistence.ensure_checkpoint_capacity(free_bytes=11, payload_bytes=10) with pytest.raises( OSError, match=r"Insufficient disk space: 0\.0 GB free, need ~0\.0 GB", ): persistence.ensure_checkpoint_capacity(free_bytes=10, payload_bytes=10) def test_checkpoint_capacity_error_reports_decimal_gigabytes_exactly(): payload_bytes = 1_000_000_000_000_000_000 with pytest.raises(OSError) as exc_info: persistence.ensure_checkpoint_capacity( free_bytes=payload_bytes, payload_bytes=payload_bytes, ) assert str(exc_info.value) == ( "Insufficient disk space: 1000000000.0 GB free, " "need ~1000000000.0 GB. " "Try a different --output-dir on a larger filesystem." ) def test_metadata_serialization_is_stable_strict_json(): first = persistence.serialize_checkpoint_metadata({"z": 1, "a": [True, None]}) second = persistence.serialize_checkpoint_metadata({"a": [True, None], "z": 1}) assert first == second assert json.loads(first) == {"a": [True, None], "z": 1} assert first == ( '{\n "a": [\n true,\n null\n ],\n "z": 1\n}' ) def test_metadata_serializer_passes_strict_format_options(monkeypatch): real_dumps = json.dumps observed = {} def record_dumps(metadata, **kwargs): observed.update(kwargs) return real_dumps(metadata, **kwargs) monkeypatch.setattr(persistence.json, "dumps", record_dumps) assert json.loads(persistence.serialize_checkpoint_metadata({"schema": 1})) == { "schema": 1, } assert observed == {"indent": 2, "sort_keys": True, "allow_nan": False} @pytest.mark.parametrize("value", [object(), float("nan"), float("inf")]) def test_metadata_serialization_rejects_nonportable_values(value): with pytest.raises((TypeError, ValueError)): persistence.serialize_checkpoint_metadata({"invalid": value}) @pytest.mark.parametrize("weights_name", ["model.safetensors", "pytorch_model.bin"]) def test_validate_local_checkpoint_accepts_complete_direct_weights(tmp_path, weights_name): _write_valid_local_checkpoint(tmp_path) if weights_name != "model.safetensors": (tmp_path / "model.safetensors").rename(tmp_path / weights_name) persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') def test_validate_local_checkpoint_accepts_complete_sharded_weights(tmp_path): metadata_json = '{"schema": 1}' _write_valid_local_checkpoint(tmp_path, metadata_json) (tmp_path / "model.safetensors").unlink() (tmp_path / "model.safetensors.index.json").write_text( json.dumps( { "metadata": {"total_size": 2}, "weight_map": { "layer.0": "model-00001-of-00002.safetensors", "layer.1": "model-00002-of-00002.safetensors", }, }, ), encoding="utf-8", ) (tmp_path / "model-00001-of-00002.safetensors").write_bytes(b"a") (tmp_path / "model-00002-of-00002.safetensors").write_bytes(b"b") persistence.validate_local_checkpoint(tmp_path, metadata_json) def test_validate_local_checkpoint_accepts_pytorch_sharded_weights(tmp_path): _write_valid_local_checkpoint(tmp_path) (tmp_path / "model.safetensors").unlink() (tmp_path / "pytorch_model.bin.index.json").write_text( json.dumps({"weight_map": {"layer": "pytorch_model-00001-of-00001.bin"}}), encoding="utf-8", ) (tmp_path / "pytorch_model-00001-of-00001.bin").write_bytes(b"weights") persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') @pytest.mark.parametrize("kind", ["file", "symlink"]) def test_validate_local_checkpoint_rejects_non_directory_staging(tmp_path, kind): staging = tmp_path / "staging" if kind == "file": staging.write_bytes(b"not a directory") else: target = tmp_path / "target" target.mkdir() staging.symlink_to(target, target_is_directory=True) with pytest.raises(ValueError, match="staging path is not a directory"): persistence.validate_local_checkpoint(staging, '{"schema": 1}') @pytest.mark.parametrize( ("mutate", "message"), [ ( lambda path: (path / "abliteration_metadata.json").write_text( "{", encoding="utf-8", ), "metadata is corrupt", ), ( lambda path: (path / "abliteration_metadata.json").write_text( '{"schema": 2}', encoding="utf-8", ), "metadata does not match", ), ( lambda path: (path / "config.json").write_text("[]", encoding="utf-8"), "model config must contain a JSON object", ), ( lambda path: (path / "tokenizer_config.json").write_bytes(b""), "tokenizer config is missing or empty", ), ( lambda path: (path / "model.safetensors").write_bytes(b""), "weights is missing or empty", ), ], ) def test_validate_local_checkpoint_rejects_corrupt_or_truncated_artifacts( tmp_path, mutate, message, ): _write_valid_local_checkpoint(tmp_path) mutate(tmp_path) with pytest.raises(ValueError, match=message): persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') def test_validate_local_checkpoint_reports_exact_metadata_mismatch(tmp_path): _write_valid_local_checkpoint(tmp_path, '{"schema": 2}') with pytest.raises(ValueError) as exc_info: persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') assert str(exc_info.value) == ( "Checkpoint metadata does not match the prepared transaction" ) @pytest.mark.parametrize( "weight_map", [ {}, {"layer": "missing.safetensors"}, {"layer": "../outside.safetensors"}, {"layer": ["not", "a", "path"]}, ], ) def test_validate_local_checkpoint_rejects_invalid_weight_indexes(tmp_path, weight_map): _write_valid_local_checkpoint(tmp_path) (tmp_path / "model.safetensors").unlink() (tmp_path / "model.safetensors.index.json").write_text( json.dumps({"weight_map": weight_map}), encoding="utf-8", ) with pytest.raises(ValueError, match="weight index|weight shard"): persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') def test_validate_local_checkpoint_reports_corrupt_weight_index_exactly(tmp_path): _write_valid_local_checkpoint(tmp_path) (tmp_path / "model.safetensors").unlink() index_path = tmp_path / "model.safetensors.index.json" index_path.write_text("{", encoding="utf-8") with pytest.raises(ValueError) as exc_info: persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') assert str(exc_info.value) == f"Checkpoint weight index is corrupt: {index_path}" def test_validate_local_checkpoint_reports_missing_weight_shard_exactly(tmp_path): _write_valid_local_checkpoint(tmp_path) (tmp_path / "model.safetensors").unlink() index_path = tmp_path / "model.safetensors.index.json" index_path.write_text( json.dumps({"weight_map": {"layer": "missing.safetensors"}}), encoding="utf-8", ) with pytest.raises(ValueError) as exc_info: persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') missing = tmp_path / "missing.safetensors" assert str(exc_info.value) == f"Checkpoint weight shard is missing or empty: {missing}" def test_validate_local_checkpoint_rejects_existing_unsafe_weight_shard(tmp_path): _write_valid_local_checkpoint(tmp_path) (tmp_path / "model.safetensors").unlink() outside = tmp_path.parent / "outside.safetensors" outside.write_bytes(b"outside") index_path = tmp_path / "model.safetensors.index.json" index_path.write_text( json.dumps({"weight_map": {"layer": "../outside.safetensors"}}), encoding="utf-8", ) with pytest.raises(ValueError) as exc_info: persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') assert str(exc_info.value) == ( f"Checkpoint weight index contains an unsafe shard path: {index_path}" ) def test_validate_local_checkpoint_rejects_checkpoint_without_weights(tmp_path): _write_valid_local_checkpoint(tmp_path) (tmp_path / "model.safetensors").unlink() with pytest.raises(ValueError, match="has no model weights"): persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') def test_validate_local_checkpoint_rejects_link_backed_weights(tmp_path): _write_valid_local_checkpoint(tmp_path) (tmp_path / "model.safetensors").unlink() outside = tmp_path / "outside.safetensors" outside.write_bytes(b"outside") (tmp_path / "model.safetensors").symlink_to(outside) with pytest.raises(ValueError, match="weights is missing or empty"): persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') def test_validate_local_checkpoint_rejects_link_backed_json(tmp_path): _write_valid_local_checkpoint(tmp_path) (tmp_path / "config.json").unlink() outside = tmp_path / "outside.json" outside.write_text("{}", encoding="utf-8") (tmp_path / "config.json").symlink_to(outside) with pytest.raises(ValueError, match="model config is missing or empty"): persistence.validate_local_checkpoint(tmp_path, '{"schema": 1}') def test_remove_checkpoint_path_handles_file_directory_symlink_and_missing(tmp_path): file_path = tmp_path / "file" file_path.write_text("data", encoding="utf-8") directory = tmp_path / "directory" directory.mkdir() (directory / "nested").write_text("data", encoding="utf-8") target = tmp_path / "target" target.mkdir() (target / "sentinel").write_text("preserve", encoding="utf-8") symlink = tmp_path / "link" symlink.symlink_to(target, target_is_directory=True) for path in (file_path, directory, symlink, tmp_path / "missing"): persistence._remove_checkpoint_path(path) assert not file_path.exists() assert not directory.exists() assert not symlink.exists() assert (target / "sentinel").read_text(encoding="utf-8") == "preserve" def test_remove_checkpoint_path_requests_race_safe_file_unlink(): path = MagicMock() path.is_symlink.return_value = True persistence._remove_checkpoint_path(path) path.unlink.assert_called_once_with(missing_ok=True) path.is_file.assert_not_called() def test_remove_checkpoint_path_requests_observable_directory_cleanup(monkeypatch): path = MagicMock() path.is_symlink.return_value = False path.is_file.return_value = False path.exists.return_value = True remove_tree = MagicMock() monkeypatch.setattr(persistence.shutil, "rmtree", remove_tree) persistence._remove_checkpoint_path(path) remove_tree.assert_called_once_with(path) def test_remove_checkpoint_path_tolerates_directory_disappearing_during_cleanup( monkeypatch, ): path = MagicMock() path.is_symlink.return_value = False path.is_file.return_value = False path.exists.return_value = True remove_tree = MagicMock(side_effect=FileNotFoundError) monkeypatch.setattr(persistence.shutil, "rmtree", remove_tree) persistence._remove_checkpoint_path(path) remove_tree.assert_called_once_with(path) def test_atomic_checkpoint_creates_parent_and_promotes_new_destination(tmp_path): destination = tmp_path / "nested" / "checkpoint" with persistence.atomic_checkpoint_directory(destination) as staging: assert staging.parent == destination.parent assert staging.name.startswith(".checkpoint.staging-") (staging / "model.bin").write_bytes(b"complete") assert (destination / "model.bin").read_bytes() == b"complete" assert list(destination.parent.glob(".checkpoint.*-*")) == [] def test_atomic_checkpoint_cleans_staging_when_new_destination_promotion_fails( tmp_path, monkeypatch, ): destination = tmp_path / "checkpoint" def fail_promotion(source, target): raise OSError(f"cannot promote {Path(source).name} to {Path(target).name}") monkeypatch.setattr(persistence.os, "replace", fail_promotion) with pytest.raises(OSError, match="cannot promote"): with persistence.atomic_checkpoint_directory(destination) as staging: (staging / "model.bin").write_bytes(b"partial") assert not destination.exists() assert list(tmp_path.glob(".checkpoint.staging-*")) == [] def test_atomic_checkpoint_replaces_symlink_without_touching_target(tmp_path): target = tmp_path / "target" target.mkdir() sentinel = target / "sentinel" sentinel.write_text("preserve", encoding="utf-8") destination = tmp_path / "checkpoint" destination.symlink_to(target, target_is_directory=True) with persistence.atomic_checkpoint_directory(destination) as staging: (staging / "model.bin").write_bytes(b"replacement") assert not destination.is_symlink() assert (destination / "model.bin").read_bytes() == b"replacement" assert sentinel.read_text(encoding="utf-8") == "preserve" def test_atomic_checkpoint_uses_os_replace_not_copy(tmp_path, monkeypatch): destination = tmp_path / "checkpoint" calls = [] real_replace = os.replace def record_replace(source, target): calls.append((Path(source), Path(target))) return real_replace(source, target) monkeypatch.setattr(persistence.os, "replace", record_replace) with persistence.atomic_checkpoint_directory(destination) as staging: (staging / "model.bin").write_bytes(b"saved") assert len(calls) == 1 assert calls[0][1] == destination assert ".staging-" in calls[0][0].name def test_atomic_checkpoint_flushes_files_before_promotion(tmp_path, monkeypatch): destination = tmp_path / "checkpoint" events = [] real_sync_file = persistence._sync_file real_replace = os.replace def record_sync(path): events.append(("sync", Path(path).name)) return real_sync_file(path) def record_replace(source, target): events.append(("replace", Path(source).name)) return real_replace(source, target) monkeypatch.setattr(persistence, "_sync_file", record_sync) monkeypatch.setattr(persistence.os, "replace", record_replace) with persistence.atomic_checkpoint_directory(destination) as staging: (staging / "model.bin").write_bytes(b"saved") assert events[0] == ("sync", "model.bin") assert events[1][0] == "replace" def test_atomic_checkpoint_flushes_nested_directories(tmp_path, monkeypatch): destination = tmp_path / "checkpoint" synced_directories = [] real_sync_directory = persistence._sync_directory def record_directory_sync(path): synced_directories.append(Path(path).name) return real_sync_directory(path) monkeypatch.setattr(persistence, "_sync_directory", record_directory_sync) with persistence.atomic_checkpoint_directory(destination) as staging: nested = staging / "nested" nested.mkdir() (nested / "payload").write_bytes(b"saved") assert "nested" in synced_directories def test_sync_directory_uses_platform_flags_and_closes_descriptor( tmp_path, monkeypatch, ): open_file = MagicMock(return_value=17) sync_file = MagicMock() close_file = MagicMock() monkeypatch.setattr(persistence, "_DIRECTORY_SYNC_FLAGS", 123) monkeypatch.setattr(persistence.os, "open", open_file) monkeypatch.setattr(persistence.os, "fsync", sync_file) monkeypatch.setattr(persistence.os, "close", close_file) persistence._sync_directory(tmp_path) open_file.assert_called_once_with(tmp_path, 123) sync_file.assert_called_once_with(17) close_file.assert_called_once_with(17) def test_sync_directory_skips_unsupported_platform(tmp_path, monkeypatch): open_file = MagicMock() monkeypatch.setattr(persistence, "_DIRECTORY_SYNC_FLAGS", None) monkeypatch.setattr(persistence.os, "open", open_file) persistence._sync_directory(tmp_path) open_file.assert_not_called() def test_sync_directory_closes_descriptor_after_fsync_failure(tmp_path, monkeypatch): close_file = MagicMock() monkeypatch.setattr(persistence, "_DIRECTORY_SYNC_FLAGS", 123) monkeypatch.setattr(persistence.os, "open", MagicMock(return_value=17)) monkeypatch.setattr( persistence.os, "fsync", MagicMock(side_effect=OSError("simulated fsync failure")), ) monkeypatch.setattr(persistence.os, "close", close_file) with pytest.raises(OSError, match="simulated fsync failure"): persistence._sync_directory(tmp_path) close_file.assert_called_once_with(17) def test_sync_checkpoint_tree_requests_bottom_up_nonfollowing_walk( tmp_path, monkeypatch, ): real_walk = os.walk observed = {} def record_walk(path, **kwargs): observed.update(kwargs) return real_walk(path, **kwargs) monkeypatch.setattr(persistence.os, "walk", record_walk) persistence._sync_checkpoint_tree(tmp_path) assert observed == {"topdown": False, "followlinks": False} @pytest.mark.parametrize("link_kind", ["file", "directory"]) def test_atomic_checkpoint_rejects_symlinks_inside_staging(tmp_path, link_kind): destination = tmp_path / "checkpoint" outside = tmp_path / "outside" if link_kind == "file": outside.write_bytes(b"outside") else: outside.mkdir() with pytest.raises(OSError, match="not a regular file|contains a symlink"): with persistence.atomic_checkpoint_directory(destination) as staging: link = staging / "link" link.symlink_to(outside, target_is_directory=link_kind == "directory") assert not destination.exists() assert outside.exists() def test_atomic_checkpoint_serializes_concurrent_complete_writers(tmp_path): destination = tmp_path / "checkpoint" destination.mkdir() (destination / "generation").write_text("old", encoding="utf-8") ready = Barrier(2) def write_generation(generation): with persistence.atomic_checkpoint_directory(destination) as staging: (staging / "generation").write_text(generation, encoding="utf-8") (staging / "payload-a").write_text(generation, encoding="utf-8") (staging / "payload-b").write_text(generation, encoding="utf-8") ready.wait(timeout=5) with ThreadPoolExecutor(max_workers=2) as pool: futures = [pool.submit(write_generation, generation) for generation in ("a", "b")] for future in futures: future.result(timeout=10) generation = (destination / "generation").read_text(encoding="utf-8") assert generation in {"a", "b"} assert (destination / "payload-a").read_text(encoding="utf-8") == generation assert (destination / "payload-b").read_text(encoding="utf-8") == generation assert list(tmp_path.glob(".checkpoint.staging-*")) == [] assert list(tmp_path.glob(".checkpoint.backup-*")) == [] def test_atomic_checkpoint_preserves_artifacts_owned_by_other_transactions(tmp_path): destination = tmp_path / "checkpoint" unrelated_staging = tmp_path / ".checkpoint.staging-operator-preserve" unrelated_backup = tmp_path / ".checkpoint.backup-operator-preserve" unrelated_staging.mkdir() unrelated_backup.mkdir() with persistence.atomic_checkpoint_directory(destination) as staging: (staging / "model.bin").write_bytes(b"saved") assert unrelated_staging.is_dir() assert unrelated_backup.is_dir() def test_rollback_rejects_ambiguous_destination_and_staging(tmp_path): destination = tmp_path / "checkpoint" staging = tmp_path / ".checkpoint.staging-owned" backup = tmp_path / ".checkpoint.backup-owned" destination.mkdir() staging.mkdir() backup.mkdir() with pytest.raises(RuntimeError) as exc_info: persistence._rollback_checkpoint_commit( destination, staging, backup, had_destination=True, ) assert str(exc_info.value) == ( "Checkpoint rollback found both destination and staging; " f"recover the previous checkpoint from {backup}" ) def test_rollback_with_no_backup_preserves_preexisting_destination(tmp_path): destination = tmp_path / "checkpoint" destination.mkdir() sentinel = destination / "sentinel" sentinel.write_text("preserve", encoding="utf-8") staging = tmp_path / ".checkpoint.staging-missing" persistence._rollback_checkpoint_commit( destination, staging, None, had_destination=True, ) assert sentinel.read_text(encoding="utf-8") == "preserve" assert not staging.exists()