152 lines
5.2 KiB
Python
152 lines
5.2 KiB
Python
import hashlib
|
|
|
|
import pytest
|
|
|
|
from tht.jobs.dwh_pipeline import DwhPreprocessPipeline
|
|
|
|
|
|
FP = "sha256:" + hashlib.sha256(b"test").hexdigest()
|
|
|
|
|
|
def test_lsh_failure_resumes_exact_run_without_repeating_introspection(tmp_path):
|
|
calls = []
|
|
|
|
def introspect(output):
|
|
calls.append("introspect")
|
|
output.write_text("catalog")
|
|
|
|
def fail_lsh(physical, output):
|
|
calls.append("lsh-failed")
|
|
raise RuntimeError("database detail that must not leak")
|
|
|
|
failed = DwhPreprocessPipeline(
|
|
workspace_id="demo",
|
|
workspace_root=tmp_path,
|
|
config_fingerprint=FP,
|
|
input_fingerprint=FP,
|
|
introspect=introspect,
|
|
build_lsh=fail_lsh,
|
|
).run(("introspect", "lsh"))
|
|
assert failed.status == "failed"
|
|
|
|
resumed = DwhPreprocessPipeline(
|
|
workspace_id="demo",
|
|
workspace_root=tmp_path,
|
|
config_fingerprint=FP,
|
|
input_fingerprint=FP,
|
|
introspect=introspect,
|
|
build_lsh=lambda physical, output: _recover_lsh(calls, output),
|
|
).run(("introspect", "lsh"), resume_run_id=failed.run_id)
|
|
|
|
assert resumed.status == "succeeded"
|
|
assert resumed.resumed_from == failed.run_id
|
|
assert calls == ["introspect", "lsh-failed", "lsh-recovered"]
|
|
|
|
|
|
def _recover_lsh(calls, output):
|
|
calls.append("lsh-recovered")
|
|
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json"):
|
|
(output / name).write_text(name)
|
|
|
|
|
|
def test_resume_rejects_a_different_stage_selection(tmp_path):
|
|
failed = DwhPreprocessPipeline(
|
|
workspace_id="demo",
|
|
workspace_root=tmp_path,
|
|
config_fingerprint=FP,
|
|
input_fingerprint=FP,
|
|
introspect=lambda output: output.write_text("catalog"),
|
|
build_lsh=lambda physical, output: (_ for _ in ()).throw(RuntimeError()),
|
|
).run(("introspect", "lsh"))
|
|
|
|
pipeline = DwhPreprocessPipeline(
|
|
workspace_id="demo",
|
|
workspace_root=tmp_path,
|
|
config_fingerprint=FP,
|
|
input_fingerprint=FP,
|
|
introspect=lambda output: output.write_text("catalog"),
|
|
build_lsh=lambda physical, output: None,
|
|
)
|
|
try:
|
|
pipeline.run(("lsh",), resume_run_id=failed.run_id)
|
|
except Exception as error:
|
|
assert "incompatible" in str(error)
|
|
else:
|
|
raise AssertionError("resume with different stages must fail")
|
|
|
|
|
|
def test_resume_rejects_tampered_succeeded_stage_artifact(tmp_path):
|
|
failed = DwhPreprocessPipeline(
|
|
workspace_id="demo", workspace_root=tmp_path,
|
|
config_fingerprint=FP, input_fingerprint=FP,
|
|
introspect=lambda output: output.write_text("catalog"),
|
|
build_lsh=lambda physical, output: (_ for _ in ()).throw(RuntimeError()),
|
|
).run(("introspect", "lsh"))
|
|
artifact = (
|
|
tmp_path / ".tht-jobs" / "dwh" / "runs" / failed.run_id
|
|
/ "artifacts" / "physical.yaml"
|
|
)
|
|
artifact.write_text("tampered")
|
|
|
|
pipeline = DwhPreprocessPipeline(
|
|
workspace_id="demo", workspace_root=tmp_path,
|
|
config_fingerprint=FP, input_fingerprint=FP,
|
|
introspect=lambda output: output.write_text("catalog"),
|
|
build_lsh=lambda physical, output: _recover_lsh([], output),
|
|
)
|
|
with pytest.raises(Exception, match="artifact manifest is invalid"):
|
|
pipeline.run(("introspect", "lsh"), resume_run_id=failed.run_id)
|
|
|
|
|
|
def test_post_publish_crash_reconciles_same_generation_on_resume(tmp_path):
|
|
crashed = False
|
|
builder_calls = 0
|
|
|
|
def crash_once(_generation):
|
|
nonlocal crashed
|
|
if not crashed:
|
|
crashed = True
|
|
raise KeyboardInterrupt("simulated process death")
|
|
|
|
def build(physical, output):
|
|
nonlocal builder_calls
|
|
builder_calls += 1
|
|
_recover_lsh([], output)
|
|
|
|
pipeline = DwhPreprocessPipeline(
|
|
workspace_id="demo", workspace_root=tmp_path,
|
|
config_fingerprint=FP, input_fingerprint=FP,
|
|
introspect=lambda output: output.write_text("catalog"),
|
|
build_lsh=build,
|
|
after_publish=crash_once,
|
|
)
|
|
with pytest.raises(KeyboardInterrupt):
|
|
pipeline.run(("introspect", "lsh"))
|
|
active = (tmp_path / ".tht-dwh" / "ACTIVE").read_text().strip()
|
|
checkpoint = next((tmp_path / ".tht-jobs" / "dwh" / "runs").glob("*/checkpoint.json"))
|
|
source_run_id = checkpoint.parent.name
|
|
assert active == source_run_id
|
|
|
|
resumed = pipeline.run(("introspect", "lsh"), resume_run_id=source_run_id)
|
|
assert resumed.status == "succeeded"
|
|
assert builder_calls == 1
|
|
assert (tmp_path / ".tht-dwh" / "ACTIVE").read_text().strip() == source_run_id
|
|
|
|
|
|
def test_resume_of_succeeded_run_detects_tampered_published_file(tmp_path):
|
|
pipeline = DwhPreprocessPipeline(
|
|
workspace_id="demo", workspace_root=tmp_path,
|
|
config_fingerprint=FP, input_fingerprint=FP,
|
|
introspect=lambda output: output.write_text("catalog"),
|
|
build_lsh=lambda physical, output: _recover_lsh([], output),
|
|
)
|
|
succeeded = pipeline.run(("introspect", "lsh"))
|
|
published = (
|
|
tmp_path / ".tht-dwh" / "generations" / succeeded.run_id / "demo_meta.json"
|
|
)
|
|
published.chmod(0o600)
|
|
published.write_text("tampered")
|
|
|
|
with pytest.raises(Exception, match="published DWH"):
|
|
pipeline.run(("introspect", "lsh"), resume_run_id=succeeded.run_id)
|