mirror of
https://github.com/elder-plinius/OBLITERATUS.git
synced 2026-08-17 16:37:30 +02:00
215 lines
7.6 KiB
Python
215 lines
7.6 KiB
Python
"""Failure and recovery tests for transactional checkpoint writes."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
from pathlib import Path
|
|
from unittest.mock import MagicMock
|
|
|
|
import pytest
|
|
import torch
|
|
|
|
import obliteratus.persistence_contracts as persistence
|
|
from obliteratus.abliterate import AbliterationPipeline, _atomic_checkpoint_directory
|
|
|
|
|
|
pytestmark = pytest.mark.cpu
|
|
|
|
|
|
def _temporary_artifacts(parent: Path, name: str) -> list[Path]:
|
|
return [
|
|
*parent.glob(f".{name}.staging-*"),
|
|
*parent.glob(f".{name}.backup-*"),
|
|
]
|
|
|
|
|
|
def test_atomic_checkpoint_replaces_existing_destination(tmp_path):
|
|
destination = tmp_path / "checkpoint"
|
|
destination.mkdir()
|
|
(destination / "old.txt").write_text("old", encoding="utf-8")
|
|
|
|
with _atomic_checkpoint_directory(destination) as staging:
|
|
(staging / "new.txt").write_text("new", encoding="utf-8")
|
|
|
|
assert not (destination / "old.txt").exists()
|
|
assert (destination / "new.txt").read_text() == "new"
|
|
assert _temporary_artifacts(tmp_path, destination.name) == []
|
|
|
|
|
|
def test_atomic_checkpoint_replaces_an_invalid_file_destination(tmp_path):
|
|
destination = tmp_path / "checkpoint"
|
|
destination.write_text("not a checkpoint", encoding="utf-8")
|
|
|
|
with _atomic_checkpoint_directory(destination) as staging:
|
|
(staging / "config.json").write_text("{}", encoding="utf-8")
|
|
|
|
assert destination.is_dir()
|
|
assert (destination / "config.json").read_text() == "{}"
|
|
assert _temporary_artifacts(tmp_path, destination.name) == []
|
|
|
|
|
|
def test_atomic_checkpoint_warns_if_obsolete_backup_cannot_be_removed(
|
|
tmp_path,
|
|
monkeypatch,
|
|
caplog,
|
|
):
|
|
destination = tmp_path / "checkpoint"
|
|
destination.write_text("old", encoding="utf-8")
|
|
real_remove = persistence._remove_checkpoint_path
|
|
|
|
def fail_backup_cleanup(path):
|
|
if ".backup-" in path.name:
|
|
raise PermissionError("simulated cleanup denial")
|
|
return real_remove(path)
|
|
|
|
monkeypatch.setattr(persistence, "_remove_checkpoint_path", fail_backup_cleanup)
|
|
with _atomic_checkpoint_directory(destination) as staging:
|
|
(staging / "config.json").write_text("{}", encoding="utf-8")
|
|
|
|
assert (destination / "config.json").is_file()
|
|
assert "could not be removed" in caplog.text
|
|
backups = list(tmp_path.glob(".checkpoint.backup-*"))
|
|
assert len(backups) == 1
|
|
assert backups[0].read_text() == "old"
|
|
|
|
|
|
def test_atomic_checkpoint_preserves_destination_on_write_failure(tmp_path):
|
|
destination = tmp_path / "checkpoint"
|
|
destination.mkdir()
|
|
sentinel = destination / "sentinel.txt"
|
|
sentinel.write_text("preserve me", encoding="utf-8")
|
|
|
|
with pytest.raises(OSError, match="simulated write failure"):
|
|
with _atomic_checkpoint_directory(destination) as staging:
|
|
(staging / "partial.bin").write_bytes(b"partial")
|
|
raise OSError("simulated write failure")
|
|
|
|
assert sentinel.read_text() == "preserve me"
|
|
assert _temporary_artifacts(tmp_path, destination.name) == []
|
|
|
|
|
|
def test_atomic_checkpoint_restores_destination_on_promotion_failure(
|
|
tmp_path,
|
|
monkeypatch,
|
|
):
|
|
destination = tmp_path / "checkpoint"
|
|
destination.mkdir()
|
|
sentinel = destination / "sentinel.txt"
|
|
sentinel.write_text("preserve me", encoding="utf-8")
|
|
real_replace = os.replace
|
|
|
|
def fail_staging_promotion(source, target):
|
|
if ".staging-" in Path(source).name:
|
|
raise OSError("simulated promotion failure")
|
|
return real_replace(source, target)
|
|
|
|
monkeypatch.setattr(persistence.os, "replace", fail_staging_promotion)
|
|
with pytest.raises(OSError, match="simulated promotion failure"):
|
|
with _atomic_checkpoint_directory(destination) as staging:
|
|
(staging / "new.txt").write_text("new", encoding="utf-8")
|
|
|
|
assert sentinel.read_text() == "preserve me"
|
|
assert _temporary_artifacts(tmp_path, destination.name) == []
|
|
|
|
|
|
def test_atomic_checkpoint_reports_recoverable_backup_when_rollback_fails(
|
|
tmp_path,
|
|
monkeypatch,
|
|
):
|
|
destination = tmp_path / "checkpoint"
|
|
destination.mkdir()
|
|
(destination / "sentinel.txt").write_text("recover me", encoding="utf-8")
|
|
real_replace = os.replace
|
|
replacement_calls = 0
|
|
|
|
def fail_promotion_and_restore(source, target):
|
|
nonlocal replacement_calls
|
|
replacement_calls += 1
|
|
if replacement_calls >= 2:
|
|
raise OSError("simulated replace failure")
|
|
return real_replace(source, target)
|
|
|
|
monkeypatch.setattr(persistence.os, "replace", fail_promotion_and_restore)
|
|
with pytest.raises(RuntimeError, match="recover the previous checkpoint from"):
|
|
with _atomic_checkpoint_directory(destination) as staging:
|
|
(staging / "new.txt").write_text("new", encoding="utf-8")
|
|
|
|
backups = list(tmp_path.glob(".checkpoint.backup-*"))
|
|
assert len(backups) == 1
|
|
assert (backups[0] / "sentinel.txt").read_text() == "recover me"
|
|
|
|
|
|
def test_rebirth_failure_preserves_checkpoint_and_owned_offload(tmp_path):
|
|
destination = tmp_path / "checkpoint"
|
|
destination.mkdir()
|
|
sentinel = destination / "sentinel.txt"
|
|
sentinel.write_text("preserve me", encoding="utf-8")
|
|
offload = tmp_path / "owned-offload"
|
|
offload.mkdir()
|
|
(offload / "weight.bin").write_bytes(b"still needed")
|
|
|
|
pipeline = AbliterationPipeline(
|
|
model_name="test-model",
|
|
output_dir=str(destination),
|
|
method="basic",
|
|
)
|
|
pipeline._on_log = lambda _message: None
|
|
pipeline._on_stage = lambda _event: None
|
|
pipeline.handle = MagicMock()
|
|
pipeline.handle.model.state_dict.return_value = {"weight": torch.ones(1)}
|
|
pipeline.handle.model.save_pretrained.side_effect = OSError("disk vanished")
|
|
pipeline.handle._offload_dir = str(offload)
|
|
pipeline.handle._owns_offload_dir = True
|
|
|
|
with pytest.raises(OSError, match="disk vanished"):
|
|
pipeline._rebirth()
|
|
|
|
assert sentinel.read_text() == "preserve me"
|
|
assert (offload / "weight.bin").read_bytes() == b"still needed"
|
|
assert _temporary_artifacts(tmp_path, destination.name) == []
|
|
|
|
|
|
def test_cleanup_only_removes_pipeline_owned_offload_directory(tmp_path):
|
|
caller_owned = tmp_path / "caller-owned"
|
|
caller_owned.mkdir()
|
|
(caller_owned / "weight.bin").write_bytes(b"owned by caller")
|
|
pipeline = AbliterationPipeline(model_name="test-model", method="basic")
|
|
pipeline._on_log = lambda _message: None
|
|
pipeline.handle = MagicMock()
|
|
pipeline.handle._offload_dir = str(caller_owned)
|
|
pipeline.handle._owns_offload_dir = False
|
|
|
|
pipeline._cleanup_offload_dir()
|
|
|
|
assert (caller_owned / "weight.bin").read_bytes() == b"owned by caller"
|
|
|
|
|
|
def test_cleanup_removes_and_clears_pipeline_owned_offload_directory(tmp_path):
|
|
owned = tmp_path / "pipeline-owned"
|
|
owned.mkdir()
|
|
(owned / "weight.bin").write_bytes(b"temporary")
|
|
pipeline = AbliterationPipeline(model_name="test-model", method="basic")
|
|
pipeline._on_log = lambda _message: None
|
|
pipeline.handle = MagicMock()
|
|
pipeline.handle._offload_dir = str(owned)
|
|
pipeline.handle._owns_offload_dir = True
|
|
|
|
pipeline._cleanup_offload_dir()
|
|
|
|
assert not owned.exists()
|
|
assert pipeline.handle._offload_dir is None
|
|
assert pipeline.handle._owns_offload_dir is False
|
|
|
|
|
|
def test_cleanup_clears_stale_owned_offload_reference(tmp_path):
|
|
pipeline = AbliterationPipeline(model_name="test-model", method="basic")
|
|
pipeline._on_log = lambda _message: None
|
|
pipeline.handle = MagicMock()
|
|
pipeline.handle._offload_dir = str(tmp_path / "already-gone")
|
|
pipeline.handle._owns_offload_dir = True
|
|
|
|
pipeline._cleanup_offload_dir()
|
|
|
|
assert pipeline.handle._offload_dir is None
|
|
assert pipeline.handle._owns_offload_dir is False
|