fix(preprocess): harden crash recovery integrity

This commit is contained in:
2026-07-12 04:49:15 +02:00
parent 1d5f8c76a7
commit c964920f16
7 changed files with 440 additions and 37 deletions
@@ -41,3 +41,33 @@ green.
One pre-existing Pydantic serialization warning is exposed by the new end-to-end job test when One pre-existing Pydantic serialization warning is exposed by the new end-to-end job test when
canonical metadata contains frozen tuple values; it does not contaminate CLI stdout. Retention is canonical metadata contains frozen tuple values; it does not contaminate CLI stdout. Retention is
an explicit stable no-op until a retention policy is configured. an explicit stable no-op until a retention policy is configured.
## Review fix wave — crash consistency and artifact integrity
Addressed all five follow-up findings:
- `JobRunner` now supports a test-only post-call/pre-checkpoint fault hook. Each stage seals a
canonical artifact manifest containing required flat filenames, SHA-256, byte size, producer
stage, and the full spec compatibility fingerprint. Resume validates the checkpoint and every
sealed artifact before allocating/copying a new run, rejecting missing, tampered, extra, nested,
or symlinked state. A sealed `running` stage is promoted after a simulated process crash; a
sealed `failed` stage is deliberately retried.
- Vector intent (exact record IDs and content hashes) is sealed before upsert. Execution reconciles
`existing_hashes` and writes only missing/mismatched rows. Crash-after-effect tests prove no
duplicate acquire, embed, or vector upsert.
- Raw upsert, stage, recovery-upsert, recovery-stage, and publish exceptions compensate the exact
generation. Compensation markers survive failed checkpoints; resume rotates the generation,
refreshes generation-bound artifacts, reconciles vectors, and stages idempotently.
- `CorpusStore.publish` is idempotent and failure-atomic. If replace succeeds but directory fsync
fails, it restores the previous `ACTIVE` value (or removes a newly created pointer), fsyncs the
rollback, and re-raises. Pipeline cleanup refuses to discard a generation referenced by ACTIVE.
- Added crash/resume coverage after all seven ordered stages; corrupt/missing plan, manifest, and
embeddings; unsafe extra paths; nonexistent run IDs; raw vector/stage failures; and post-replace
ACTIVE rollback.
Fresh fix-wave verification:
- Focused jobs/corpus/CLI/search suite: `82 passed, 17 warnings`.
- Available harness suite (same sandbox exclusions described above):
`579 passed, 5 deselected, 31 warnings`.
- Scoped Ruff and `git diff --check`: clean.
+136
View File
@@ -49,6 +49,11 @@ class Vectors:
raise RuntimeError("partial write") raise RuntimeError("partial write")
return len(records) 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): def delete_generation(self, collection, generation):
self.records = [ self.records = [
value for value in 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 embedder.calls == []
assert vectors.records == [] assert vectors.records == []
assert result.generation is None and result.published is False 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,
)
+15
View File
@@ -41,3 +41,18 @@ def test_active_pointer_cannot_escape_generation_root(tmp_path):
store.active_path.write_text("../outside\n") store.active_path.write_text("../outside\n")
with pytest.raises(UnsafeCorpusPath): with pytest.raises(UnsafeCorpusPath):
store.active_manifest() 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
+61 -1
View File
@@ -4,7 +4,7 @@ import pytest
from pydantic import ValidationError from pydantic import ValidationError
from tht.jobs.models import JobSpec 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 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 = context.run_dir / "artifacts"
artifacts.mkdir() artifacts.mkdir()
(artifacts / "discovery.json").write_text('{"source":"one"}') (artifacts / "discovery.json").write_text('{"source":"one"}')
return StageArtifacts(("discovery.json",))
first = run_job( first = run_job(
_spec(tmp_path, stage_ids=("discover", "acquire")), _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" 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): def test_successful_job_is_idempotently_resumable(tmp_path):
calls = [] calls = []
+68 -24
View File
@@ -16,7 +16,7 @@ from tht.ports.evidence import EvidenceSource, SourceObject
from tht.ports.vector import VectorStore, VectorWriteRecord from tht.ports.vector import VectorStore, VectorWriteRecord
from tht.vectorstore.records import VectorRecord from tht.vectorstore.records import VectorRecord
from tht.jobs.models import JobSpec from tht.jobs.models import JobSpec
from tht.jobs.runner import JobContext, run_job from tht.jobs.runner import JobContext, StageArtifacts, run_job, seal_stage_artifacts
EVIDENCE_STAGE_IDS = ( EVIDENCE_STAGE_IDS = (
@@ -96,6 +96,7 @@ class CorpusPipeline:
input_fingerprint: str, input_fingerprint: str,
dry_run: bool = False, dry_run: bool = False,
resume_run_id: str | None = None, resume_run_id: str | None = None,
after_stage_return=None,
) -> PipelineResult: ) -> PipelineResult:
"""Execute preprocessing through the durable shared job envelope.""" """Execute preprocessing through the durable shared job envelope."""
discovered = self._discover() discovered = self._discover()
@@ -159,10 +160,11 @@ class CorpusPipeline:
"removed": removed, "removed": removed,
"previous": previous.model_dump(mode="json") if previous else None, "previous": previous.model_dump(mode="json") if previous else None,
}) })
return StageArtifacts(("plan.json",))
def acquire_stage(context: JobContext) -> None: def acquire_stage(context: JobContext) -> None:
if context.dry_run: if context.dry_run:
return return StageArtifacts()
plan = read(context, "plan.json") plan = read(context, "plan.json")
previous = CorpusManifest.model_validate(plan["previous"]) if plan["previous"] else None previous = CorpusManifest.model_validate(plan["previous"]) if plan["previous"] else None
prior = {doc.source_id: doc for doc in previous.documents} if previous else {} prior = {doc.source_id: doc for doc in previous.documents} if previous else {}
@@ -194,10 +196,11 @@ class CorpusPipeline:
}, },
) )
write(context, "manifest.json", manifest.model_dump(mode="json")) write(context, "manifest.json", manifest.model_dump(mode="json"))
return StageArtifacts(("manifest.json",))
def embed_stage(context: JobContext) -> None: def embed_stage(context: JobContext) -> None:
if context.dry_run: if context.dry_run:
return return StageArtifacts()
plan = read(context, "plan.json") plan = read(context, "plan.json")
manifest = CorpusManifest.model_validate(read(context, "manifest.json")) manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
changed_docs = {doc.document_id for doc in manifest.documents if doc.source_id in plan["changed"]} changed_docs = {doc.document_id for doc in manifest.documents if doc.source_id in plan["changed"]}
@@ -208,6 +211,7 @@ class CorpusPipeline:
): ):
raise PipelineError("embedding output is incompatible") raise PipelineError("embedding output is incompatible")
write(context, "embeddings.json", embeddings) write(context, "embeddings.json", embeddings)
return StageArtifacts(("embeddings.json",))
def records(context: JobContext): def records(context: JobContext):
plan = read(context, "plan.json") plan = read(context, "plan.json")
@@ -220,7 +224,8 @@ class CorpusPipeline:
def compensate(context: JobContext) -> None: def compensate(context: JobContext) -> None:
generation = read(context, "plan.json")["generation"] generation = read(context, "plan.json")["generation"]
self.store.discard(generation) if self.store.active_generation() != generation:
self.store.discard(generation)
try: try:
self.vector_store.delete_generation("evidence", generation) self.vector_store.delete_generation("evidence", generation)
except Exception: except Exception:
@@ -241,62 +246,101 @@ class CorpusPipeline:
for document in manifest.documents: for document in manifest.documents:
if document.source_id in changed and generations.get(document.document_id) == old: if document.source_id in changed and generations.get(document.document_id) == old:
generations[document.document_id] = plan["generation"] generations[document.document_id] = plan["generation"]
metadata = dict(manifest.metadata) manifest_payload = manifest.model_dump(mode="json")
metadata["document_generations"] = generations manifest_payload["metadata"]["document_generations"] = generations
manifest = manifest.model_copy(update={ manifest_payload["vector_generation"] = plan["generation"]
"vector_generation": plan["generation"], "metadata": metadata, manifest = CorpusManifest.model_validate(manifest_payload)
})
write(context, "manifest.json", manifest.model_dump(mode="json")) write(context, "manifest.json", manifest.model_dump(mode="json"))
marker.unlink() marker.unlink()
def vector_stage(context: JobContext) -> None: def vector_stage(context: JobContext) -> None:
if context.dry_run: if context.dry_run:
return return StageArtifacts()
rotate_compensated_generation(context) rotate_compensated_generation(context)
values = records(context) values = records(context)
write(context, "vector-intent.json", {
"generation": read(context, "plan.json")["generation"],
"records": {value.record.id: value.content_hash for value in values},
})
seal_stage_artifacts(
context, "vector_upsert",
("plan.json", "manifest.json", "vector-intent.json"), spec,
)
try: try:
if values and self.vector_store.upsert("evidence", values) != len(values): existing = self.vector_store.existing_hashes("evidence", ["evidence"])
missing = [
value for value in values
if existing.get(value.record.id) != value.content_hash
]
if missing and self.vector_store.upsert("evidence", missing) != len(missing):
raise PipelineError("vector write count mismatch") raise PipelineError("vector write count mismatch")
except Exception: except Exception:
compensate(context) compensate(context)
raise raise
return StageArtifacts(("plan.json", "manifest.json", "vector-intent.json"))
def stage_stage(context: JobContext) -> None: def stage_stage(context: JobContext) -> None:
if context.dry_run: if context.dry_run:
return return StageArtifacts()
plan = read(context, "plan.json") plan = read(context, "plan.json")
manifest = CorpusManifest.model_validate(read(context, "manifest.json")) manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
recovered = False
try: try:
self.store.stage( if artifact(context, "compensated.json").exists():
manifest, {doc.document_id: doc.content for doc in manifest.documents}, recovered = True
generation=plan["generation"], rotate_compensated_generation(context)
) values = records(context)
existing = self.vector_store.existing_hashes("evidence", ["evidence"])
missing = [value for value in values if existing.get(value.record.id) != value.content_hash]
if missing and self.vector_store.upsert("evidence", missing) != len(missing):
raise PipelineError("vector write count mismatch")
plan = read(context, "plan.json")
manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
if not self.store.generation_path(plan["generation"]).exists():
self.store.stage(
manifest, {doc.document_id: doc.content for doc in manifest.documents},
generation=plan["generation"],
)
self.store.manifest(plan["generation"]) self.store.manifest(plan["generation"])
except Exception: except Exception:
compensate(context) compensate(context)
raise raise
return StageArtifacts(
("plan.json", "manifest.json", "vector-intent.json") if recovered else ()
)
def publish_stage(context: JobContext) -> None: def publish_stage(context: JobContext) -> None:
if context.dry_run: if context.dry_run:
return return StageArtifacts()
if artifact(context, "compensated.json").exists(): if artifact(context, "compensated.json").exists():
rotate_compensated_generation(context) rotate_compensated_generation(context)
values = records(context) values = records(context)
if values and self.vector_store.upsert("evidence", values) != len(values): try:
existing = self.vector_store.existing_hashes("evidence", ["evidence"])
missing = [value for value in values if existing.get(value.record.id) != value.content_hash]
if missing and self.vector_store.upsert("evidence", missing) != len(missing):
raise PipelineError("vector write count mismatch")
except Exception:
compensate(context) compensate(context)
raise PipelineError("vector write count mismatch") raise
manifest = CorpusManifest.model_validate(read(context, "manifest.json")) manifest = CorpusManifest.model_validate(read(context, "manifest.json"))
generation = read(context, "plan.json")["generation"] generation = read(context, "plan.json")["generation"]
self.store.stage( try:
manifest, {doc.document_id: doc.content for doc in manifest.documents}, if not self.store.generation_path(generation).exists():
generation=generation, self.store.stage(
) manifest, {doc.document_id: doc.content for doc in manifest.documents},
generation=generation,
)
except Exception:
compensate(context)
raise
generation = read(context, "plan.json")["generation"] generation = read(context, "plan.json")["generation"]
try: try:
self.store.publish(generation) self.store.publish(generation)
except Exception: except Exception:
compensate(context) compensate(context)
raise raise
return StageArtifacts(("plan.json", "manifest.json", "vector-intent.json"))
def retention_stage(context: JobContext) -> None: def retention_stage(context: JobContext) -> None:
# Retention policy is intentionally a stable no-op until configured. # Retention policy is intentionally a stable no-op until configured.
@@ -305,7 +349,7 @@ class CorpusPipeline:
report = run_job(spec, [ report = run_job(spec, [
discover_stage, acquire_stage, embed_stage, vector_stage, discover_stage, acquire_stage, embed_stage, vector_stage,
stage_stage, publish_stage, retention_stage, stage_stage, publish_stage, retention_stage,
]) ], after_stage_return=after_stage_return)
run_dir = workspace_root / ".tht-jobs" / "evidence" / "runs" / report.run_id run_dir = workspace_root / ".tht-jobs" / "evidence" / "runs" / report.run_id
plan = json.loads((run_dir / "artifacts" / "plan.json").read_text()) plan = json.loads((run_dir / "artifacts" / "plan.json").read_text())
if dry_run: if dry_run:
+24 -3
View File
@@ -46,6 +46,7 @@ class CorpusStore:
self.root = Path(root) self.root = Path(root)
self.active_path = self.root / "ACTIVE" self.active_path = self.root / "ACTIVE"
self._replace = os.replace self._replace = os.replace
self._fsync_directory = self._sync_root
self._ensure_root() self._ensure_root()
def _ensure_root(self) -> None: def _ensure_root(self) -> None:
@@ -110,15 +111,35 @@ class CorpusStore:
manifest = self.manifest(generation) manifest = self.manifest(generation)
if manifest.manifest_id != generation: if manifest.manifest_id != generation:
raise UnsafeCorpusPath("manifest generation mismatch") raise UnsafeCorpusPath("manifest generation mismatch")
if self.active_generation() == generation:
return generation
previous = self.active_generation()
temporary = self.active_path.with_name(f".ACTIVE.{uuid.uuid4().hex}.tmp") temporary = self.active_path.with_name(f".ACTIVE.{uuid.uuid4().hex}.tmp")
_atomic_write(temporary, (generation + "\n").encode()) replaced = False
self._replace(temporary, self.active_path) try:
_atomic_write(temporary, (generation + "\n").encode())
self._replace(temporary, self.active_path)
replaced = True
self._fsync_directory()
except BaseException:
temporary.unlink(missing_ok=True)
if replaced:
if previous is None:
self.active_path.unlink(missing_ok=True)
else:
rollback = self.active_path.with_name(f".ACTIVE.rollback.{uuid.uuid4().hex}.tmp")
_atomic_write(rollback, (previous + "\n").encode())
self._replace(rollback, self.active_path)
self._sync_root()
raise
return generation
def _sync_root(self) -> None:
directory = os.open(self.root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW) directory = os.open(self.root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
try: try:
os.fsync(directory) os.fsync(directory)
finally: finally:
os.close(directory) os.close(directory)
return generation
def active_generation(self) -> str | None: def active_generation(self) -> str | None:
try: try:
+106 -9
View File
@@ -23,6 +23,11 @@ class CorruptCheckpointError(RuntimeError):
pass pass
@dataclass(frozen=True)
class StageArtifacts:
required: tuple[str, ...] = ()
@dataclass(frozen=True) @dataclass(frozen=True)
class JobContext: class JobContext:
run_id: str run_id: str
@@ -35,6 +40,84 @@ class JobContext:
Stage = Callable[[JobContext], Any] Stage = Callable[[JobContext], Any]
def _artifact_digest(path: Path) -> dict[str, Any]:
payload = path.read_bytes()
return {"sha256": hashlib.sha256(payload).hexdigest(), "size": len(payload)}
def _seal_artifacts(context: JobContext, stage: str, result: Any, spec: JobSpec) -> None:
required = result.required if isinstance(result, StageArtifacts) else ()
root = context.run_dir / "artifacts"
root.mkdir(exist_ok=True)
manifest_path = root / "artifact-manifest.json"
manifest = json.loads(manifest_path.read_text()) if manifest_path.exists() else {
"schema_version": 1,
"spec_fingerprint": _compatibility_fingerprint(spec, list(spec.stage_ids)),
"stages": {},
}
files = {}
for relative in required:
candidate = root / relative
if Path(relative).is_absolute() or ".." in Path(relative).parts or candidate.is_symlink():
raise CorruptCheckpointError("artifact path is unsafe")
if not candidate.is_file():
raise CorruptCheckpointError("required stage artifact is missing")
files[relative] = _artifact_digest(candidate)
for prior in manifest["stages"].values():
if relative in prior.get("required", []):
prior["required"].remove(relative)
prior["files"].pop(relative, None)
manifest["stages"][stage] = {"required": list(required), "files": files}
_atomic_write(manifest_path, json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n")
def seal_stage_artifacts(
context: JobContext, stage: str, required: tuple[str, ...], spec: JobSpec,
) -> None:
"""Durably record external-effect intent before a stage performs that effect."""
_seal_artifacts(context, stage, StageArtifacts(required), spec)
def _validate_artifacts(run_dir: Path, spec: JobSpec, source: JobRun) -> set[str]:
root = run_dir / "artifacts"
manifest_path = root / "artifact-manifest.json"
successful = {stage.name for stage in source.stages if stage.status == "succeeded"}
if not successful and not manifest_path.exists():
return set()
try:
manifest = json.loads(manifest_path.read_text())
if manifest["spec_fingerprint"] != _compatibility_fingerprint(spec, list(spec.stage_ids)):
raise ValueError
sealed = set(manifest["stages"])
allowed = {"artifact-manifest.json"}
for stage, record in manifest["stages"].items():
for relative in record["required"]:
candidate = root / relative
if Path(relative).is_absolute() or ".." in Path(relative).parts or candidate.is_symlink():
raise ValueError
if not candidate.is_file() or _artifact_digest(candidate) != record["files"][relative]:
raise ValueError
allowed.add(relative)
if not successful.issubset(sealed):
raise ValueError
entries = list(root.iterdir())
if any(path.is_symlink() or not path.is_file() for path in entries):
raise ValueError
actual = {path.name for path in entries}
incomplete = next((stage for stage in source.stages if stage.status != "succeeded"), None)
marker = root / "compensated.json"
if incomplete is not None and incomplete.status in {"failed", "running"} and marker.is_file():
payload = json.loads(marker.read_text())
if not isinstance(payload.get("generation"), str):
raise ValueError
allowed.add("compensated.json")
if actual != allowed:
raise ValueError
return sealed
except (OSError, KeyError, TypeError, ValueError, json.JSONDecodeError) as error:
raise CorruptCheckpointError("resume artifact manifest is invalid") from error
def _atomic_write(path: Path, payload: str) -> None: def _atomic_write(path: Path, payload: str) -> None:
temporary = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp") temporary = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600) fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
@@ -91,7 +174,10 @@ def _new_run(spec: JobSpec, run_id: str, stages: Sequence[Stage]) -> JobRun:
) )
def _resume_run(spec: JobSpec, run_id: str, stages: Sequence[Stage], source: JobRun) -> JobRun: def _resume_run(
spec: JobSpec, run_id: str, stages: Sequence[Stage], source: JobRun,
effect_completed: set[str] | None = None,
) -> JobRun:
requested_names = list(spec.stage_ids) requested_names = list(spec.stage_ids)
if len(requested_names) != len(stages) or len(requested_names) != len(set(requested_names)): if len(requested_names) != len(stages) or len(requested_names) != len(set(requested_names)):
raise CorruptCheckpointError("resume checkpoint is incompatible with requested stages") raise CorruptCheckpointError("resume checkpoint is incompatible with requested stages")
@@ -107,11 +193,15 @@ def _resume_run(spec: JobSpec, run_id: str, stages: Sequence[Stage], source: Job
resumed_stages = [] resumed_stages = []
for name in requested_names: for name in requested_names:
previous = source_by_name.get(name) previous = source_by_name.get(name)
resumed_stages.append( if previous is not None and previous.status == "succeeded":
previous resumed_stages.append(previous)
if previous is not None and previous.status == "succeeded" elif previous is not None and previous.status == "running" and name in (effect_completed or set()):
else StageRun(name=name) resumed_stages.append(StageRun(
) name=name, status="succeeded", started_at=previous.started_at or utc_now(),
finished_at=utc_now(),
))
else:
resumed_stages.append(StageRun(name=name))
return JobRun( return JobRun(
run_id=run_id, run_id=run_id,
compatibility_fingerprint=source.compatibility_fingerprint, compatibility_fingerprint=source.compatibility_fingerprint,
@@ -129,7 +219,10 @@ def _resume_run(spec: JobSpec, run_id: str, stages: Sequence[Stage], source: Job
) )
def run_job(spec: JobSpec, stages: Sequence[Stage]) -> JobReport: def run_job(
spec: JobSpec, stages: Sequence[Stage], *,
after_stage_return: Callable[[JobContext, str], Any] | None = None,
) -> JobReport:
"""Run stages once, returning a terminal report instead of leaking stage exceptions.""" """Run stages once, returning a terminal report instead of leaking stage exceptions."""
with WorkspaceJobLock(spec.workspace_root, spec.workspace_id, spec.job_type): with WorkspaceJobLock(spec.workspace_root, spec.workspace_id, spec.job_type):
jobs_root = spec.workspace_root / ".tht-jobs" / spec.job_type / "runs" jobs_root = spec.workspace_root / ".tht-jobs" / spec.job_type / "runs"
@@ -147,6 +240,7 @@ def run_job(spec: JobSpec, stages: Sequence[Stage]) -> JobReport:
source_path = matches[0] source_path = matches[0]
source = _load_checkpoint(source_path) source = _load_checkpoint(source_path)
_validate_resume_source(spec, stages, source) _validate_resume_source(spec, stages, source)
effect_completed = _validate_artifacts(source_path.parent, spec, source)
run_id = uuid.uuid4().hex run_id = uuid.uuid4().hex
run_dir = jobs_root / run_id run_dir = jobs_root / run_id
@@ -155,7 +249,7 @@ def run_job(spec: JobSpec, stages: Sequence[Stage]) -> JobReport:
if source is None: if source is None:
run = _new_run(spec, run_id, stages) run = _new_run(spec, run_id, stages)
else: else:
run = _resume_run(spec, run_id, stages, source) run = _resume_run(spec, run_id, stages, source, effect_completed)
source_artifacts = jobs_root / source.run_id / "artifacts" source_artifacts = jobs_root / source.run_id / "artifacts"
if source_artifacts.exists(): if source_artifacts.exists():
shutil.copytree(source_artifacts, run_dir / "artifacts") shutil.copytree(source_artifacts, run_dir / "artifacts")
@@ -173,7 +267,7 @@ def run_job(spec: JobSpec, stages: Sequence[Stage]) -> JobReport:
) )
_persist(checkpoint_path, run) _persist(checkpoint_path, run)
try: try:
stage_callable(context) stage_result = stage_callable(context)
except Exception: except Exception:
failed = stage.model_copy( failed = stage.model_copy(
update={ update={
@@ -191,6 +285,9 @@ def run_job(spec: JobSpec, stages: Sequence[Stage]) -> JobReport:
) )
_persist(checkpoint_path, run) _persist(checkpoint_path, run)
break break
_seal_artifacts(context, stage.name, stage_result, spec)
if after_stage_return is not None:
after_stage_return(context, stage.name)
succeeded = stage.model_copy(update={"status": "succeeded", "finished_at": utc_now()}) succeeded = stage.model_copy(update={"status": "succeeded", "finished_at": utc_now()})
run = run.model_copy( run = run.model_copy(
update={"stages": run.stages[:index] + (succeeded,) + run.stages[index + 1 :]} update={"stages": run.stages[:index] + (succeeded,) + run.stages[index + 1 :]}