fix(preprocess): reconcile durable DWH generations

This commit is contained in:
2026-07-12 05:51:49 +02:00
parent 05accc1443
commit 78ff360882
7 changed files with 295 additions and 30 deletions
+90 -1
View File
@@ -6,6 +6,7 @@ from typer.testing import CliRunner
from tht.cli import app
from tht.jobs.dwh_pipeline import DwhPreprocessPipeline
from tht.jobs.dwh_pipeline import resolve_dwh_snapshot
from tht.jobs.locking import _lock_name
@@ -36,7 +37,8 @@ def test_selected_dwh_stages_run_in_declared_order(tmp_path):
published = tmp_path / ".tht-dwh" / "generations" / active
assert (published / "physical.yaml").read_text() == "catalog"
assert sorted(path.name for path in published.iterdir()) == [
"demo_lsh.pkl", "demo_meta.json", "demo_minhashes.pkl", "physical.yaml"
"demo_lsh.pkl", "demo_meta.json", "demo_minhashes.pkl",
"generation-manifest.json", "physical.yaml",
]
@@ -139,3 +141,90 @@ def test_failed_multi_file_build_never_replaces_active_generation(tmp_path):
],
).run(("introspect", "lsh"), resume_run_id=failed.run_id)
assert resumed.status == "succeeded"
def test_unsafe_lsh_filename_is_rejected(tmp_path):
import pytest
with pytest.raises(ValueError, match="flat safe"):
DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: None, build_lsh=lambda physical, output: None,
lsh_filenames=("../escape.pkl", "ok.pkl", "meta.json"),
)
def test_active_fsync_failure_restores_previous_pointer(monkeypatch, tmp_path):
def build(physical, output):
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json"):
(output / name).write_text(name)
first_pipeline = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("old"), build_lsh=build,
)
first = first_pipeline.run()
original_fsync = first_pipeline._fsync
failed_once = False
def fail_active_once(path):
nonlocal failed_once
if path.name == ".tht-dwh" and not failed_once:
failed_once = True
raise OSError("injected directory fsync failure")
original_fsync(path)
second = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text("new"), build_lsh=build,
)
monkeypatch.setattr(second, "_fsync", fail_active_once)
failed = second.run()
assert failed.status == "failed"
assert (tmp_path / ".tht-dwh" / "ACTIVE").read_text().strip() == first.run_id
def test_snapshot_stays_on_one_generation_across_publish(tmp_path):
from types import SimpleNamespace
def pipeline(content):
return DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: output.write_text(content),
build_lsh=lambda physical, output: [
(output / name).write_text(content)
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
)
first = pipeline("old").run()
cfg = SimpleNamespace(paths=SimpleNamespace(
artifacts=tmp_path / "artifacts", indexes=tmp_path / "indexes"
))
snapshot = resolve_dwh_snapshot(cfg)
pipeline("new").run()
assert snapshot.generation == first.run_id
assert snapshot.physical.read_text() == "old"
assert (snapshot.lsh_dir / "demo_meta.json").read_text() == "old"
def test_generation_retention_keeps_active_and_one_rollback(tmp_path):
run_ids = []
for index in range(5):
report = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output, i=index: output.write_text(str(i)),
build_lsh=lambda physical, output, i=index: [
(output / name).write_text(str(i))
for name in ("demo_lsh.pkl", "demo_minhashes.pkl", "demo_meta.json")
],
retain_generations=2,
).run()
run_ids.append(report.run_id)
remaining = {path.name for path in (tmp_path / ".tht-dwh" / "generations").iterdir()}
assert remaining == set(run_ids[-2:])
+11 -3
View File
@@ -100,6 +100,7 @@ def test_resume_rejects_tampered_succeeded_stage_artifact(tmp_path):
def test_post_publish_crash_reconciles_same_generation_on_resume(tmp_path):
crashed = False
builder_calls = 0
def crash_once(_generation):
nonlocal crashed
@@ -107,11 +108,16 @@ def test_post_publish_crash_reconciles_same_generation_on_resume(tmp_path):
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=lambda physical, output: _recover_lsh([], output),
build_lsh=build,
after_publish=crash_once,
)
with pytest.raises(KeyboardInterrupt):
@@ -123,7 +129,8 @@ def test_post_publish_crash_reconciles_same_generation_on_resume(tmp_path):
resumed = pipeline.run(("introspect", "lsh"), resume_run_id=source_run_id)
assert resumed.status == "succeeded"
assert (tmp_path / ".tht-dwh" / "ACTIVE").read_text().strip() == resumed.run_id
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):
@@ -137,7 +144,8 @@ def test_resume_of_succeeded_run_detects_tampered_published_file(tmp_path):
published = (
tmp_path / ".tht-dwh" / "generations" / succeeded.run_id / "demo_meta.json"
)
published.chmod(0o600)
published.write_text("tampered")
with pytest.raises(Exception, match="digest mismatch"):
with pytest.raises(Exception, match="published DWH"):
pipeline.run(("introspect", "lsh"), resume_run_id=succeeded.run_id)
+2 -5
View File
@@ -9,12 +9,9 @@ lsh_app = typer.Typer(help="Indice LSH su valori dei campi (derivato, rigenerabi
def _lsh_dir(cfg) -> Path:
from tht.jobs.dwh_pipeline import active_generation_dir
from tht.jobs.dwh_pipeline import resolve_dwh_snapshot
active = active_generation_dir(cfg.paths.artifacts.parent)
if active is not None:
return active
return cfg.paths.indexes / "lsh"
return resolve_dwh_snapshot(cfg).lsh_dir
def _extract_lsh_values(dwh, physical, annotations, limit):
+2 -5
View File
@@ -37,12 +37,9 @@ def _load_config_or_exit(config: Path):
def physical_path(cfg) -> Path:
from tht.jobs.dwh_pipeline import active_generation_dir
from tht.jobs.dwh_pipeline import resolve_dwh_snapshot
active = active_generation_dir(cfg.paths.artifacts.parent)
if active is not None:
return active / "physical.yaml"
return cfg.paths.artifacts / "mschema" / "physical.yaml"
return resolve_dwh_snapshot(cfg).physical
def annotations_path(cfg) -> Path:
+10 -6
View File
@@ -47,6 +47,9 @@ def search_cmd(
from tht.search import combined_search
cfg = _load_config_or_exit(config)
from tht.jobs.dwh_pipeline import resolve_dwh_snapshot
dwh_snapshot = resolve_dwh_snapshot(cfg)
require_vector_cfg(cfg)
from tht.search.evidence import active_searcher
@@ -88,7 +91,7 @@ def search_cmd(
lsh_hits = None
try:
lsh, minhashes, meta = load_index(
cfg.paths.indexes / "lsh", name=cfg.database.db_schema
dwh_snapshot.lsh_dir, name=cfg.database.db_schema
)
hits = query_index(lsh, minhashes, keyword, meta, top_n=top * 3)
lsh_hits = [(h.table, h.column, h.value, h.score) for h in hits]
@@ -100,12 +103,12 @@ def search_cmd(
)
if kind == "schema":
from tht.cli.schema_cmd import annotations_path, physical_path
from tht.cli.schema_cmd import annotations_path
from tht.mschema.models import Annotations, PhysicalSchema
from tht.mschema.render import to_mschema_text
from tht.search import schema_tables
phys_file = physical_path(cfg)
phys_file = dwh_snapshot.physical
if not phys_file.exists():
typer.secho(
f"ERRORE: {phys_file} non trovato. Esegui prima `tht schema introspect`.",
@@ -241,6 +244,9 @@ def pack_cmd(
from tht.vectorstore.rest_client import VectorRestError
cfg = _load_config_or_exit(config)
from tht.jobs.dwh_pipeline import resolve_dwh_snapshot
dwh_snapshot = resolve_dwh_snapshot(cfg)
require_vector_cfg(cfg)
tables: list[dict] = []
@@ -261,10 +267,8 @@ def pack_cmd(
warnings.append(f"retrieval non disponibile ({e}): prosegui con le ricerche live")
if vec is not None:
from tht.cli.schema_cmd import physical_path
descriptions: dict[str, str] = {}
phys_file = physical_path(cfg)
phys_file = dwh_snapshot.physical
if phys_file.exists():
from tht.mschema.models import PhysicalSchema
+177 -10
View File
@@ -10,6 +10,7 @@ import shutil
import stat
import uuid
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
from tht.jobs.models import JobReport, JobSpec
@@ -24,15 +25,84 @@ from tht.jobs.runner import (
DWH_STAGE_IDS = ("introspect", "lsh")
_RUN_ID = re.compile(r"^[0-9a-f]{32}$")
_SAFE_FILE = re.compile(r"^[A-Za-z0-9_-]+\.(?:pkl|json)$")
GENERATION_MANIFEST = "generation-manifest.json"
@dataclass(frozen=True)
class DwhArtifactSnapshot:
generation: str | None
physical: Path
lsh_dir: Path
def _digest(path: Path) -> str:
return hashlib.sha256(path.read_bytes()).hexdigest()
def _read_owned(path: Path, *, readonly: bool) -> bytes:
fd = os.open(path, os.O_RDONLY | os.O_NOFOLLOW)
try:
info = os.fstat(fd)
if (
not stat.S_ISREG(info.st_mode)
or info.st_uid != os.getuid()
or info.st_nlink != 1
or (readonly and bool(info.st_mode & stat.S_IWUSR))
):
raise OSError("unsafe DWH generation file")
chunks = []
while chunk := os.read(fd, 1024 * 1024):
chunks.append(chunk)
return b"".join(chunks)
finally:
os.close(fd)
def validate_generation(target: Path) -> dict:
manifest_path = target / GENERATION_MANIFEST
try:
manifest = json.loads(_read_owned(manifest_path, readonly=True).decode("utf-8"))
files = manifest["files"]
if (
manifest["generation"] != target.name
or not isinstance(files, dict)
or not re.fullmatch(r"sha256:[0-9a-f]{64}", manifest["job_spec_fingerprint"])
or not re.fullmatch(r"[0-9a-f]{64}", manifest["artifact_manifest_sha256"])
):
raise ValueError
if set(path.name for path in target.iterdir()) != set(files) | {GENERATION_MANIFEST}:
raise ValueError
for name, expected in files.items():
if name != "physical.yaml" and not _SAFE_FILE.fullmatch(name):
raise ValueError
path = target / name
if hashlib.sha256(_read_owned(path, readonly=True)).hexdigest() != expected:
raise ValueError
return manifest
except (OSError, KeyError, TypeError, ValueError, UnicodeDecodeError, json.JSONDecodeError) as error:
raise CorruptCheckpointError("published DWH generation is invalid") from error
def resolve_dwh_snapshot(cfg) -> DwhArtifactSnapshot:
target = active_generation_dir(cfg.paths.artifacts.parent)
if target is None:
return DwhArtifactSnapshot(
None, cfg.paths.artifacts / "mschema" / "physical.yaml", cfg.paths.indexes / "lsh"
)
validate_generation(target)
return DwhArtifactSnapshot(target.name, target / "physical.yaml", target)
def active_generation_dir(workspace_root: Path) -> Path | None:
pointer = workspace_root / ".tht-dwh" / "ACTIVE"
try:
generation = pointer.read_text(encoding="utf-8").strip()
generation = _read_owned(pointer, readonly=False).decode("utf-8").strip()
except FileNotFoundError:
return None
if not _RUN_ID.fullmatch(generation) or pointer.is_symlink():
except (OSError, UnicodeDecodeError) as error:
raise CorruptCheckpointError("DWH ACTIVE pointer is invalid") from error
if not _RUN_ID.fullmatch(generation):
raise CorruptCheckpointError("DWH ACTIVE pointer is invalid")
target = pointer.parent / "generations" / generation
if not target.is_dir() or target.is_symlink():
@@ -56,6 +126,7 @@ class DwhPreprocessPipeline:
current_physical: Path | None = None,
current_lsh_dir: Path | None = None,
after_publish: Callable[[str], object] | None = None,
retain_generations: int = 3,
) -> None:
self.workspace_id = workspace_id
self.workspace_root = workspace_root
@@ -70,6 +141,13 @@ class DwhPreprocessPipeline:
self.current_physical = current_physical
self.current_lsh_dir = current_lsh_dir
self.after_publish = after_publish
if isinstance(retain_generations, bool) or retain_generations < 1:
raise ValueError("retain_generations must be positive")
self.retain_generations = retain_generations
if len(set(self.lsh_filenames)) != 3 or any(
not _SAFE_FILE.fullmatch(name) for name in self.lsh_filenames
):
raise ValueError("LSH filenames must be unique flat safe names")
def run(
self, steps: tuple[str, ...] = DWH_STAGE_IDS, *, resume_run_id: str | None = None
@@ -119,7 +197,25 @@ class DwhPreprocessPipeline:
return self._publish_stage(context, "lsh", spec)
implementations = {"introspect": introspect_stage, "lsh": lsh_stage}
return run_job(spec, tuple(implementations[step] for step in steps))
return run_job(
spec, tuple(implementations[step] for step in steps),
reconcile_effects=self._reconcile_effects,
)
def _reconcile_effects(self, source, run_dir: Path) -> set[str]:
running = next(
(stage for stage in source.stages if stage.status == "running" and stage.effect_state == "intent"),
None,
)
if running is None:
return set()
target = self.workspace_root / ".tht-dwh" / "generations" / source.run_id
validate_generation(target)
active = active_generation_dir(self.workspace_root)
if active != target:
raise CorruptCheckpointError("sealed DWH publication is not ACTIVE")
self._validate_published(target, run_dir / "artifacts", running.artifact_files)
return {running.name}
def _validate_resume_publication(self, run_id: str) -> None:
run_dir = self.workspace_root / ".tht-jobs" / "dwh" / "runs" / run_id
@@ -180,6 +276,7 @@ class DwhPreprocessPipeline:
self._publish(context.run_id, artifacts, required)
if self.after_publish is not None:
self.after_publish(context.run_id)
self._cleanup_generations()
return StageArtifacts(required)
def _publish(self, generation: str, artifacts: Path, required: tuple[str, ...]) -> None:
@@ -199,12 +296,37 @@ class DwhPreprocessPipeline:
shutil.copyfile(artifacts / name, destination)
with destination.open("rb") as stream:
os.fsync(stream.fileno())
destination.chmod(0o400)
manifest = {
"schema_version": 1,
"generation": generation,
"files": {name: _digest(artifacts / name) for name in required},
"job_spec_fingerprint": json.loads(
(artifacts / "artifact-manifest.json").read_text(encoding="utf-8")
)["spec_fingerprint"],
"artifact_manifest_sha256": _digest(
artifacts / "artifact-manifest.json"
),
}
manifest_path = temporary / GENERATION_MANIFEST
manifest_path.write_text(
json.dumps(manifest, sort_keys=True, separators=(",", ":")) + "\n",
encoding="utf-8",
)
manifest_path.chmod(0o400)
with manifest_path.open("rb") as stream:
os.fsync(stream.fileno())
self._fsync(temporary)
os.replace(temporary, target)
self._fsync(generations)
except BaseException:
shutil.rmtree(temporary, ignore_errors=True)
raise
pointer = root / "ACTIVE"
try:
previous = pointer.read_text(encoding="utf-8")
except FileNotFoundError:
previous = None
pointer_tmp = root / f".ACTIVE.{uuid.uuid4().hex}.tmp"
fd = os.open(pointer_tmp, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
try:
@@ -212,20 +334,25 @@ class DwhPreprocessPipeline:
stream.write(generation + "\n")
stream.flush()
os.fsync(stream.fileno())
os.replace(pointer_tmp, root / "ACTIVE")
self._fsync(root)
os.replace(pointer_tmp, pointer)
try:
self._fsync(root)
except BaseException:
self._restore_pointer(root, pointer, previous)
raise
except BaseException:
pointer_tmp.unlink(missing_ok=True)
raise
@staticmethod
def _validate_published(target: Path, artifacts: Path, required: tuple[str, ...]) -> None:
if (
target.is_symlink()
or not target.is_dir()
or {path.name for path in target.iterdir()} != set(required)
):
manifest = validate_generation(target)
if set(manifest["files"]) != set(required):
raise CorruptCheckpointError("published DWH generation is invalid")
if manifest["artifact_manifest_sha256"] != _digest(
artifacts / "artifact-manifest.json"
):
raise CorruptCheckpointError("published DWH job manifest digest mismatch")
for name in required:
source, published = artifacts / name, target / name
if published.is_symlink() or not published.is_file():
@@ -235,6 +362,46 @@ class DwhPreprocessPipeline:
).digest():
raise CorruptCheckpointError("published DWH artifact digest mismatch")
def _restore_pointer(self, root: Path, pointer: Path, previous: str | None) -> None:
if previous is None:
pointer.unlink(missing_ok=True)
else:
restore = root / f".ACTIVE.restore.{uuid.uuid4().hex}.tmp"
restore.write_text(previous, encoding="utf-8")
with restore.open("rb") as stream:
os.fsync(stream.fileno())
os.replace(restore, pointer)
self._fsync(root)
def _cleanup_generations(self) -> None:
root = self.workspace_root / ".tht-dwh" / "generations"
if not root.exists():
return
active = active_generation_dir(self.workspace_root)
active_name = active.name if active else None
protected = {active_name} if active_name else set()
runs = self.workspace_root / ".tht-jobs" / "dwh" / "runs"
for checkpoint in runs.glob("*/checkpoint.json") if runs.exists() else ():
try:
status = json.loads(checkpoint.read_text(encoding="utf-8"))["status"]
if status in {"running", "failed"}:
protected.add(checkpoint.parent.name)
except (OSError, KeyError, ValueError):
continue
generations = sorted(
(path for path in root.iterdir() if path.is_dir() and _RUN_ID.fullmatch(path.name)),
key=lambda path: path.stat().st_mtime_ns,
)
keep_recent = {path.name for path in generations[-self.retain_generations:]}
for path in generations:
if path.name in protected | keep_recent:
continue
path.chmod(0o700)
for child in path.iterdir():
child.chmod(0o600)
shutil.rmtree(path)
self._fsync(root)
@staticmethod
def _ensure_owned_dir(path: Path) -> None:
try:
+3
View File
@@ -249,6 +249,7 @@ def _resume_run(
def run_job(
spec: JobSpec, stages: Sequence[Stage], *,
after_stage_return: Callable[[JobContext, str], Any] | None = None,
reconcile_effects: Callable[[JobRun, Path], set[str]] | None = None,
) -> JobReport:
"""Run stages once, returning a terminal report instead of leaking stage exceptions."""
with WorkspaceJobLock(spec.workspace_root, spec.workspace_id, spec.job_type):
@@ -268,6 +269,8 @@ def run_job(
source = _load_checkpoint(source_path)
_validate_resume_source(spec, stages, source)
effect_completed = _validate_artifacts(source_path.parent, spec, source)
if reconcile_effects is not None:
effect_completed |= reconcile_effects(source, source_path.parent)
run_id = uuid.uuid4().hex
run_dir = jobs_root / run_id