Files
ThothII/harness/tests/test_job_runner.py

401 lines
14 KiB
Python

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