Files
OBLITERATUS/tests/test_auto_obliterate.py
T

266 lines
7.8 KiB
Python

"""Behavior tests for resumable auto-obliteration orchestration."""
from __future__ import annotations
import json
from pathlib import Path
import pytest
import obliteratus.abliterate
from obliteratus.auto_obliterate import (
AutoObliterateResult,
AutoObliterator,
IterationResult,
)
pytestmark = pytest.mark.cpu
def _finish(generator):
yielded = []
while True:
try:
yielded.append(next(generator))
except StopIteration as completed:
return yielded, completed.value
class _SuccessfulPipeline:
created = []
def __init__(self, **kwargs):
self.kwargs = kwargs
self._quality_metrics = {
"perplexity": 4.25,
"coherence": 0.9,
"refusal_rate": 0.04,
"kl_divergence": 0.02,
}
self._strong_layers = [0, 1]
self._expert_directions = {0: {0: object(), 1: object()}}
self.handle = type("Handle", (), {"model": object(), "tokenizer": object()})()
self.created.append(self)
def run(self):
Path(self.kwargs["output_dir"]).mkdir(parents=True)
self.kwargs["on_log"]("pipeline ran")
return Path(self.kwargs["output_dir"])
def test_result_round_trip_ignores_unknown_forward_compatible_fields():
original = AutoObliterateResult(
model_id="local/model",
iterations=[
IterationResult(
iteration=1,
method="aggressive",
prompt_volume=8,
categories_targeted=["test"],
),
],
success=True,
)
encoded = original.to_dict()
encoded["future_field"] = "ignored"
encoded["iterations"][0]["future_field"] = "ignored"
restored = AutoObliterateResult.from_dict(encoded)
assert restored == original
assert isinstance(restored.iterations[0], IterationResult)
def test_valid_state_resumes_at_next_iteration(tmp_path):
output = tmp_path / "state"
output.mkdir()
state = AutoObliterateResult(
model_id="local/model",
iterations=[IterationResult(1, "aggressive", 8, output_dir="first")],
)
(output / "auto_state.json").write_text(json.dumps(state.to_dict()))
auto = AutoObliterator("local/model", output_base=str(output), max_iterations=2)
assert auto._resume_from == 1
assert auto._result.iterations[0].method == "aggressive"
@pytest.mark.parametrize(
"invalid_state",
[
"not json",
"[]",
'{"model_id": "different/model"}',
'{"model_id": "local/model", "iterations": [5]}',
],
)
def test_invalid_state_is_quarantined_with_actionable_warning(
tmp_path,
caplog,
invalid_state,
):
output = tmp_path / "state"
output.mkdir()
state_file = output / "auto_state.json"
state_file.write_text(invalid_state)
auto = AutoObliterator("local/model", output_base=str(output))
quarantined = list(output.glob("auto_state.json.corrupt-*"))
assert auto._resume_from == 0
assert not state_file.exists()
assert len(quarantined) == 1
assert quarantined[0].read_text() == invalid_state
assert "quarantined at" in caplog.text
def test_invalid_state_is_retained_when_quarantine_fails(tmp_path, monkeypatch, caplog):
output = tmp_path / "state"
output.mkdir()
state_file = output / "auto_state.json"
state_file.write_text("not json", encoding="utf-8")
def fail_quarantine(_source, _target):
raise OSError("read-only filesystem")
monkeypatch.setattr("obliteratus.auto_obliterate.os.replace", fail_quarantine)
auto = AutoObliterator("local/model", output_base=str(output))
assert auto._resume_from == 0
assert state_file.read_text() == "not json"
assert "could not be quarantined" in caplog.text
assert "read-only filesystem" in caplog.text
def test_interrupted_state_replace_preserves_previous_checkpoint(
tmp_path,
monkeypatch,
caplog,
):
output = tmp_path / "state"
output.mkdir()
state_file = output / "auto_state.json"
state_file.write_text('{"sentinel": true}', encoding="utf-8")
auto = AutoObliterator.__new__(AutoObliterator)
auto.model_id = "local/model"
auto.output_base = str(output)
auto._state_file = state_file
auto._result = AutoObliterateResult(model_id="local/model")
def fail_replace(_self, _target):
raise OSError("simulated interrupted replace")
monkeypatch.setattr(Path, "replace", fail_replace)
auto._save_state()
assert json.loads(state_file.read_text()) == {"sentinel": True}
assert not (output / "auto_state.tmp").exists()
assert "simulated interrupted replace" in caplog.text
def test_auto_loop_runs_pipeline_persists_metrics_and_stops_at_target(
tmp_path,
monkeypatch,
):
_SuccessfulPipeline.created.clear()
monkeypatch.setattr(
obliteratus.abliterate,
"AbliterationPipeline",
_SuccessfulPipeline,
)
monkeypatch.setattr(
AutoObliterator,
"_get_expanded_prompts",
staticmethod(lambda _iteration: (["harmful"] * 3, ["harmless"] * 3)),
)
auto = AutoObliterator(
"local/model",
output_base=str(tmp_path / "run"),
max_iterations=3,
target_refusal_rate=0.05,
trust_remote_code=False,
)
yielded, result = _finish(auto.run())
assert result.success is True
assert result.final_refusal_rate == 0.04
assert len(result.iterations) == 1
assert result.iterations[0].strong_layers == 2
assert result.iterations[0].ega_expert_dirs == 2
assert result.final_output_dir.endswith("iter_1")
assert len(yielded) == 5
assert yielded[-1][0] == "✅ Complete"
created = _SuccessfulPipeline.created[0]
assert created.kwargs["model_name"] == "local/model"
assert created.kwargs["method"] == "aggressive"
assert created.kwargs["harmful_prompts"] == ["harmful"] * 3
saved = json.loads((tmp_path / "run" / "auto_state.json").read_text())
assert saved["success"] is True
def test_auto_loop_records_failures_and_completes_without_success(
tmp_path,
monkeypatch,
):
class FailingPipeline:
def __init__(self, **_kwargs):
pass
def run(self):
raise RuntimeError("simulated pipeline failure")
monkeypatch.setattr(obliteratus.abliterate, "AbliterationPipeline", FailingPipeline)
monkeypatch.setattr(
AutoObliterator,
"_get_expanded_prompts",
staticmethod(lambda _iteration: (["harmful"], ["harmless"])),
)
auto = AutoObliterator(
"local/model",
output_base=str(tmp_path / "run"),
max_iterations=1,
)
yielded, result = _finish(auto.run())
assert result.success is False
assert result.final_output_dir == ""
assert result.iterations[0].error == "simulated pipeline failure"
assert yielded[-1][0] == "⚠️ Complete (target not met)"
assert "simulated pipeline" in yielded[-1][2]
def test_prompt_expansion_and_benchmark_fallbacks(monkeypatch):
monkeypatch.delenv("OPENROUTER_API_KEY", raising=False)
assert AutoObliterator._quick_benchmark_claude("missing", "model") == {
"method": "skipped",
"reason": "no OPENROUTER_API_KEY",
}
pipeline = type(
"Pipeline",
(),
{"_quality_metrics": {"refusal_rate": 0.2, "coherence": 0.8}},
)()
assert AutoObliterator._quick_benchmark_heuristic(pipeline) == {
"refusal_rate": 0.2,
"perplexity": None,
"coherence": 0.8,
"kl_divergence": None,
"method": "heuristic",
}
def test_reset_clears_persisted_state(tmp_path):
auto = AutoObliterator("local/model", output_base=str(tmp_path / "run"))
auto._result.iterations.append(IterationResult(1, "aggressive", 1))
auto._save_state()
auto.reset()
assert auto._resume_from == 0
assert auto._result.iterations == []
assert not auto._state_file.exists()