fix(preprocess): harden crash recovery integrity
This commit is contained in:
@@ -49,6 +49,11 @@ class Vectors:
|
||||
raise RuntimeError("partial write")
|
||||
return len(records)
|
||||
|
||||
def existing_hashes(self, collection, kinds):
|
||||
return {
|
||||
value.record.id: value.content_hash for value in self.records
|
||||
}
|
||||
|
||||
def delete_generation(self, collection, generation):
|
||||
self.records = [
|
||||
value for value in self.records
|
||||
@@ -173,3 +178,134 @@ def test_job_pipeline_dry_run_only_discovers_and_reports_changes(tmp_path):
|
||||
assert embedder.calls == []
|
||||
assert vectors.records == []
|
||||
assert result.generation is None and result.published is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("crash_stage", [
|
||||
"discover", "acquire_normalize_chunk", "embed", "vector_upsert",
|
||||
"stage_validate", "publish", "retention_cleanup",
|
||||
])
|
||||
def test_job_pipeline_crash_after_each_stage_resumes_without_duplicate_effects(tmp_path, crash_stage):
|
||||
one = item("one", "a")
|
||||
source = Source([(one, "hello")])
|
||||
embedder = Embedder()
|
||||
vectors = Vectors()
|
||||
candidate = pipeline(tmp_path, source, embedder=embedder, vectors=vectors)
|
||||
|
||||
class Crash(BaseException):
|
||||
pass
|
||||
|
||||
def fault(_context, stage):
|
||||
if stage == crash_stage:
|
||||
raise Crash()
|
||||
|
||||
with pytest.raises(Crash):
|
||||
candidate.run_as_job(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint="sha256:" + "1" * 64,
|
||||
input_fingerprint="sha256:" + "2" * 64,
|
||||
after_stage_return=fault,
|
||||
)
|
||||
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
|
||||
crashed = next(runs.iterdir()).name
|
||||
result = candidate.run_as_job(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint="sha256:" + "1" * 64,
|
||||
input_fingerprint="sha256:" + "2" * 64,
|
||||
resume_run_id=crashed,
|
||||
)
|
||||
assert result.status == "succeeded"
|
||||
assert source.acquire_calls == ["fs:one"]
|
||||
assert len(embedder.calls) == 1
|
||||
assert len(vectors.records) == 1
|
||||
|
||||
|
||||
def test_job_pipeline_raw_upsert_failure_compensates_and_resumes_with_new_generation(tmp_path):
|
||||
one = item("one", "a")
|
||||
vectors = Vectors(fail=True)
|
||||
candidate = pipeline(tmp_path, Source([(one, "hello")]), vectors=vectors)
|
||||
first = candidate.run_as_job(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint="sha256:" + "1" * 64,
|
||||
input_fingerprint="sha256:" + "2" * 64,
|
||||
)
|
||||
assert first.status == "failed"
|
||||
assert vectors.records == []
|
||||
old_generation = first.generation
|
||||
vectors.fail = False
|
||||
resumed = candidate.run_as_job(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint="sha256:" + "1" * 64,
|
||||
input_fingerprint="sha256:" + "2" * 64,
|
||||
resume_run_id=first.run_id,
|
||||
)
|
||||
assert resumed.status == "succeeded", resumed
|
||||
assert resumed.generation != old_generation
|
||||
assert candidate.store.active_generation() == resumed.generation
|
||||
|
||||
|
||||
def test_job_pipeline_raw_stage_failure_compensates_vectors_and_resumes(tmp_path, monkeypatch):
|
||||
one = item("one", "a")
|
||||
vectors = Vectors()
|
||||
candidate = pipeline(tmp_path, Source([(one, "hello")]), vectors=vectors)
|
||||
real_stage = candidate.store.stage
|
||||
calls = 0
|
||||
|
||||
def fail_once(*args, **kwargs):
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
if calls == 1:
|
||||
raise OSError("raw stage failure")
|
||||
return real_stage(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(candidate.store, "stage", fail_once)
|
||||
first = candidate.run_as_job(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint="sha256:" + "1" * 64,
|
||||
input_fingerprint="sha256:" + "2" * 64,
|
||||
)
|
||||
assert first.status == "failed" and vectors.records == []
|
||||
resumed = candidate.run_as_job(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint="sha256:" + "1" * 64,
|
||||
input_fingerprint="sha256:" + "2" * 64,
|
||||
resume_run_id=first.run_id,
|
||||
)
|
||||
assert resumed.status == "succeeded", resumed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stage,filename", [
|
||||
("discover", "plan.json"),
|
||||
("acquire_normalize_chunk", "manifest.json"),
|
||||
("embed", "embeddings.json"),
|
||||
])
|
||||
@pytest.mark.parametrize("mutation", ["missing", "tampered"])
|
||||
def test_job_pipeline_rejects_corrupt_required_artifacts_before_resume(
|
||||
tmp_path, stage, filename, mutation,
|
||||
):
|
||||
one = item("one", "a")
|
||||
candidate = pipeline(tmp_path, Source([(one, "hello")]))
|
||||
|
||||
class Crash(BaseException):
|
||||
pass
|
||||
|
||||
with pytest.raises(Crash):
|
||||
candidate.run_as_job(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint="sha256:" + "1" * 64,
|
||||
input_fingerprint="sha256:" + "2" * 64,
|
||||
after_stage_return=lambda _context, name: (
|
||||
(_ for _ in ()).throw(Crash()) if name == stage else None
|
||||
),
|
||||
)
|
||||
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
|
||||
crashed = next(runs.iterdir())
|
||||
target = crashed / "artifacts" / filename
|
||||
target.unlink() if mutation == "missing" else target.write_text("tampered")
|
||||
from tht.jobs.runner import CorruptCheckpointError
|
||||
with pytest.raises(CorruptCheckpointError, match="artifact"):
|
||||
candidate.run_as_job(
|
||||
workspace_id="demo", workspace_root=tmp_path,
|
||||
config_fingerprint="sha256:" + "1" * 64,
|
||||
input_fingerprint="sha256:" + "2" * 64,
|
||||
resume_run_id=crashed.name,
|
||||
)
|
||||
|
||||
@@ -41,3 +41,18 @@ def test_active_pointer_cannot_escape_generation_root(tmp_path):
|
||||
store.active_path.write_text("../outside\n")
|
||||
with pytest.raises(UnsafeCorpusPath):
|
||||
store.active_manifest()
|
||||
|
||||
|
||||
def test_publish_restores_previous_active_when_directory_fsync_fails_after_replace(tmp_path, monkeypatch):
|
||||
store = CorpusStore(tmp_path / "corpus")
|
||||
first = store.stage(CorpusManifest(), {})
|
||||
second = store.stage(CorpusManifest(), {})
|
||||
store.publish(first)
|
||||
def fail_once():
|
||||
store._fsync_directory = store._sync_root
|
||||
raise OSError("post replace crash")
|
||||
|
||||
store._fsync_directory = fail_once
|
||||
with pytest.raises(OSError, match="post replace"):
|
||||
store.publish(second)
|
||||
assert store.active_generation() == first
|
||||
|
||||
@@ -4,7 +4,7 @@ import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from tht.jobs.models import JobSpec
|
||||
from tht.jobs.runner import CorruptCheckpointError, run_job
|
||||
from tht.jobs.runner import CorruptCheckpointError, StageArtifacts, run_job
|
||||
import tht.jobs.runner as runner_module
|
||||
|
||||
|
||||
@@ -69,6 +69,7 @@ def test_resume_carries_successful_stage_artifacts_into_new_run(tmp_path):
|
||||
artifacts = context.run_dir / "artifacts"
|
||||
artifacts.mkdir()
|
||||
(artifacts / "discovery.json").write_text('{"source":"one"}')
|
||||
return StageArtifacts(("discovery.json",))
|
||||
|
||||
first = run_job(
|
||||
_spec(tmp_path, stage_ids=("discover", "acquire")),
|
||||
@@ -85,6 +86,65 @@ def test_resume_carries_successful_stage_artifacts_into_new_run(tmp_path):
|
||||
assert resumed.status == "succeeded"
|
||||
|
||||
|
||||
def test_crash_after_stage_effect_resumes_without_repeating_stage(tmp_path):
|
||||
calls = []
|
||||
|
||||
def stage(context):
|
||||
calls.append("stage")
|
||||
artifacts = context.run_dir / "artifacts"
|
||||
artifacts.mkdir()
|
||||
(artifacts / "effect.json").write_text("ok")
|
||||
return StageArtifacts(("effect.json",))
|
||||
|
||||
class Crash(BaseException):
|
||||
pass
|
||||
|
||||
with pytest.raises(Crash):
|
||||
run_job(
|
||||
_spec(tmp_path), [stage],
|
||||
after_stage_return=lambda *_: (_ for _ in ()).throw(Crash()),
|
||||
)
|
||||
runs = tmp_path / ".tht-jobs" / "evidence" / "runs"
|
||||
crashed_run = next(runs.iterdir()).name
|
||||
resumed = run_job(_spec(tmp_path).with_resume(crashed_run), [stage])
|
||||
assert resumed.status == "succeeded"
|
||||
assert calls == ["stage"]
|
||||
|
||||
|
||||
def test_resume_rejects_tampered_successful_stage_artifact(tmp_path):
|
||||
def stage(context):
|
||||
artifacts = context.run_dir / "artifacts"
|
||||
artifacts.mkdir()
|
||||
(artifacts / "effect.json").write_text("ok")
|
||||
return StageArtifacts(("effect.json",))
|
||||
|
||||
report = run_job(_spec(tmp_path), [stage])
|
||||
path = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id / "artifacts" / "effect.json"
|
||||
path.write_text("tampered")
|
||||
with pytest.raises(CorruptCheckpointError, match="artifact"):
|
||||
run_job(_spec(tmp_path).with_resume(report.run_id), [stage])
|
||||
|
||||
|
||||
def test_resume_rejects_extra_symlink_before_any_stage(tmp_path):
|
||||
report = run_job(_spec(tmp_path), [lambda _context: StageArtifacts()])
|
||||
artifacts = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id / "artifacts"
|
||||
(artifacts / "unsafe").symlink_to(tmp_path)
|
||||
called = False
|
||||
|
||||
def forbidden(_context):
|
||||
nonlocal called
|
||||
called = True
|
||||
|
||||
with pytest.raises(CorruptCheckpointError, match="artifact"):
|
||||
run_job(_spec(tmp_path).with_resume(report.run_id), [forbidden])
|
||||
assert called is False
|
||||
|
||||
|
||||
def test_nonexistent_well_formed_resume_run_id_is_rejected(tmp_path):
|
||||
with pytest.raises(CorruptCheckpointError, match="checkpoint"):
|
||||
run_job(_spec(tmp_path).with_resume("a" * 32), [lambda _context: None])
|
||||
|
||||
|
||||
def test_successful_job_is_idempotently_resumable(tmp_path):
|
||||
calls = []
|
||||
|
||||
|
||||
Reference in New Issue
Block a user