import json import os import pytest from pydantic import ValidationError from tht.jobs.models import JobSpec from tht.jobs.runner import CorruptCheckpointError, StageArtifacts, run_job import tht.jobs.runner as runner_module def _spec(tmp_path, **updates): values = { "workspace_id": "demo", "job_type": "evidence", "workspace_root": tmp_path, "spec_version": "jobs-v1", "pipeline_version": "evidence-v1", "config_fingerprint": "sha256:" + "1" * 64, "input_fingerprint": "sha256:" + "2" * 64, "stage_ids": ("stage",), } values.update(updates) return JobSpec(**values) def test_job_models_are_immutable(tmp_path): spec = _spec(tmp_path) with pytest.raises(ValidationError, match="Instance is frozen"): spec.job_type = "dwh" resumed = spec.with_resume("a" * 32) assert resumed.workspace_root == tmp_path assert resumed.resume_run_id == "a" * 32 def test_failed_stage_is_resumable_and_skips_completed_stage(tmp_path): calls = [] def discover(context): calls.append(("discover", context.dry_run)) def acquire(_context): calls.append(("acquire", False)) raise RuntimeError("source /customer/alice token=secret unavailable") first = run_job(_spec(tmp_path, stage_ids=("discover", "acquire")), [discover, acquire]) assert first.status == "failed" assert [stage.status for stage in first.stages] == ["succeeded", "failed"] assert first.stages[1].error.model_dump() == { "category": "internal", "code": "stage_exception", "message": "stage execution failed", } def acquire(_context): calls.append(("recovered", False)) second = run_job( _spec(tmp_path, resume_run_id=first.run_id, stage_ids=("discover", "acquire")), [discover, acquire], ) assert second.resumed_from == first.run_id assert second.status == "succeeded" assert calls == [("discover", False), ("acquire", False), ("recovered", False)] def test_resume_carries_successful_stage_artifacts_into_new_run(tmp_path): def discover(context): 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")), [discover, lambda _context: (_ for _ in ()).throw(RuntimeError("crash"))], ) def acquire(context): assert (context.run_dir / "artifacts" / "discovery.json").read_text() == '{"source":"one"}' resumed = run_job( _spec(tmp_path, resume_run_id=first.run_id, stage_ids=("discover", "acquire")), [discover, acquire], ) 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]) @pytest.mark.parametrize("tamper", ["artifact_and_manifest", "spec", "producer"]) def test_resume_rejects_manifest_root_or_binding_tamper(tmp_path, tamper): 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]) artifacts = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id / "artifacts" manifest_path = artifacts / "artifact-manifest.json" manifest = json.loads(manifest_path.read_text()) if tamper == "artifact_and_manifest": (artifacts / "effect.json").write_text("evil") digest = __import__("hashlib").sha256(b"evil").hexdigest() manifest["stages"]["stage"]["files"]["effect.json"] = { "sha256": digest, "size": 4, } elif tamper == "spec": manifest["spec_fingerprint"] = "sha256:" + "0" * 64 else: manifest["stages"]["other"] = manifest["stages"].pop("stage") manifest_path.write_text(json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n") with pytest.raises(CorruptCheckpointError, match="artifact"): run_job(_spec(tmp_path).with_resume(report.run_id), [stage]) def test_successful_job_is_idempotently_resumable(tmp_path): calls = [] def normalize(_context): calls.append("normalize") first = run_job(_spec(tmp_path), [normalize]) second = run_job(_spec(tmp_path, resume_run_id=first.run_id), [normalize]) assert first.status == second.status == "succeeded" assert second.resumed_from == first.run_id assert calls == ["normalize"] def test_checkpoints_and_report_are_json_safe_and_do_not_disclose_workspace_path(tmp_path): def publish(_context): return {"ignored": "/customer/alice", "password": "secret"} report = run_job(_spec(tmp_path), [publish]) run_dir = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id checkpoint = json.loads((run_dir / "checkpoint.json").read_text()) payload = (run_dir / "report.json").read_text() parsed = json.loads(payload) assert checkpoint["status"] == "succeeded" assert parsed["schema_version"] == 1 assert parsed["run_id"] == report.run_id assert str(tmp_path) not in payload assert "alice" not in payload assert "secret" not in payload assert parsed["started_at"].endswith("Z") assert parsed["finished_at"].endswith("Z") def test_dry_run_is_exposed_to_stages_and_report(tmp_path): observed = [] def plan(context): observed.append(context.dry_run) report = run_job(_spec(tmp_path, dry_run=True), [plan]) assert observed == [True] assert report.dry_run is True assert report.status == "succeeded" def test_corrupt_checkpoint_is_rejected_without_running_stages(tmp_path): first = run_job(_spec(tmp_path), [lambda _context: None]) checkpoint = ( tmp_path / ".tht-jobs" / "evidence" / "runs" / first.run_id / "checkpoint.json" ) checkpoint.write_text("{not-json") called = False def stage(_context): nonlocal called called = True with pytest.raises(CorruptCheckpointError, match="checkpoint is invalid"): run_job(_spec(tmp_path, resume_run_id=first.run_id), [stage]) assert called is False def test_stage_timestamps_are_aware_and_ordered(tmp_path): report = run_job(_spec(tmp_path), [lambda _context: None]) stage = report.stages[0] assert stage.started_at.tzinfo is not None assert stage.finished_at.tzinfo is not None assert stage.started_at <= stage.finished_at assert report.started_at <= stage.started_at <= report.finished_at @pytest.mark.parametrize( ("update", "replacement"), [ ("dry_run", True), ("spec_version", "jobs-v2"), ("pipeline_version", "evidence-v2"), ("config_fingerprint", "sha256:" + "3" * 64), ("input_fingerprint", "sha256:" + "4" * 64), ("workspace_id", "other"), ("job_type", "dwh"), ], ) def test_resume_rejects_changed_identity_or_inputs_before_stage_execution( tmp_path, update, replacement ): first = run_job(_spec(tmp_path), [lambda _context: None]) called = False def stage(_context): nonlocal called called = True values = {update: replacement, "resume_run_id": first.run_id} with pytest.raises(CorruptCheckpointError, match="incompatible"): run_job(_spec(tmp_path, **values), [stage]) assert called is False @pytest.mark.parametrize("stages", [[], [lambda _context: None, lambda _context: None]]) def test_resume_rejects_removed_or_inserted_stages(tmp_path, stages): def first_stage(_context): pass first = run_job(_spec(tmp_path), [first_stage]) with pytest.raises(CorruptCheckpointError, match="incompatible"): run_job(_spec(tmp_path, resume_run_id=first.run_id), stages) def test_resume_rejects_reordered_stages(tmp_path): def one(_context): pass def two(_context): pass first = run_job(_spec(tmp_path, stage_ids=("one", "two")), [one, two]) with pytest.raises(CorruptCheckpointError, match="incompatible"): run_job( _spec(tmp_path, resume_run_id=first.run_id, stage_ids=("two", "one")), [two, one], ) def test_hostile_exception_identity_never_enters_terminal_report(tmp_path): Hostile = type("ApiKey_secret_/customer/alice", (Exception,), {}) def fail(_context): raise Hostile("password=hunter2") report = run_job(_spec(tmp_path), [fail]) payload = report.model_dump_json() assert report.status == "failed" assert report.stages[0].error.model_dump() == { "category": "internal", "code": "stage_exception", "message": "stage execution failed", } assert "secret" not in payload assert "alice" not in payload assert "hunter2" not in payload def test_new_run_without_resume_allows_intentional_spec_change(tmp_path): first = run_job(_spec(tmp_path), [lambda _context: None]) second = run_job(_spec(tmp_path, input_fingerprint="sha256:" + "9" * 64), [lambda _context: None]) assert second.status == "succeeded" assert second.run_id != first.run_id assert second.resumed_from is None def test_run_directories_are_private_and_fsynced_before_atomic_replace(tmp_path, monkeypatch): events = [] real_replace = os.replace monkeypatch.setattr(runner_module.os, "fsync", lambda _fd: events.append("fsync")) def tracked_replace(source, destination): events.append("replace") real_replace(source, destination) monkeypatch.setattr(runner_module.os, "replace", tracked_replace) report = run_job(_spec(tmp_path), [lambda _context: None]) run_dir = tmp_path / ".tht-jobs" / "evidence" / "runs" / report.run_id assert run_dir.stat().st_mode & 0o777 == 0o700 first_replace = events.index("replace") assert "fsync" in events[:first_replace] assert "fsync" in events[first_replace + 1 :] def _tamper_checkpoint(tmp_path, report, transform): path = tmp_path / ".tht-jobs" / report.job_type / "runs" / report.run_id / "checkpoint.json" payload = json.loads(path.read_text()) transform(payload) path.write_text(json.dumps(payload)) @pytest.mark.parametrize( "transform", [ lambda payload: payload["stages"].pop(), lambda payload: payload["stages"].reverse(), lambda payload: payload["stages"].append(dict(payload["stages"][0])), lambda payload: payload["stages"][0].update(name="substitute"), lambda payload: payload.update(input_fingerprint="sha256:" + "f" * 64), lambda payload: payload["stages"][0].update(status="pending"), ], ) def test_semantically_tampered_checkpoint_fails_without_orphan_run(tmp_path, transform): def one(_context): pass def two(_context): pass spec = _spec(tmp_path, stage_ids=("one", "two")) first = run_job(spec, [one, two]) runs = tmp_path / ".tht-jobs" / "evidence" / "runs" before = {path.name for path in runs.iterdir()} _tamper_checkpoint(tmp_path, first, transform) called = False def forbidden(_context): nonlocal called called = True with pytest.raises(CorruptCheckpointError, match="invalid|incompatible"): run_job(spec.with_resume(first.run_id), [forbidden, forbidden]) assert called is False assert {path.name for path in runs.iterdir()} == before def test_stored_compatibility_fingerprint_tamper_fails_without_orphan(tmp_path): first = run_job(_spec(tmp_path), [lambda _context: None]) runs = tmp_path / ".tht-jobs" / "evidence" / "runs" before = {path.name for path in runs.iterdir()} _tamper_checkpoint( tmp_path, first, lambda payload: payload.update(compatibility_fingerprint="sha256:" + "0" * 64), ) with pytest.raises(CorruptCheckpointError, match="invalid"): run_job(_spec(tmp_path, resume_run_id=first.run_id), [lambda _context: None]) assert {path.name for path in runs.iterdir()} == before