diff --git a/harness/tests/test_dwh_preprocess_job.py b/harness/tests/test_dwh_preprocess_job.py index d612e750..aa86b1a3 100644 --- a/harness/tests/test_dwh_preprocess_job.py +++ b/harness/tests/test_dwh_preprocess_job.py @@ -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:]) diff --git a/harness/tests/test_lsh_job_resume.py b/harness/tests/test_lsh_job_resume.py index 8b3f3368..81d9c4a0 100644 --- a/harness/tests/test_lsh_job_resume.py +++ b/harness/tests/test_lsh_job_resume.py @@ -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) diff --git a/harness/tht/cli/lsh_cmd.py b/harness/tht/cli/lsh_cmd.py index e341d1df..b4e62439 100644 --- a/harness/tht/cli/lsh_cmd.py +++ b/harness/tht/cli/lsh_cmd.py @@ -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): diff --git a/harness/tht/cli/schema_cmd.py b/harness/tht/cli/schema_cmd.py index 92471422..04dcbbb9 100644 --- a/harness/tht/cli/schema_cmd.py +++ b/harness/tht/cli/schema_cmd.py @@ -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: diff --git a/harness/tht/cli/search_cmd.py b/harness/tht/cli/search_cmd.py index 2766e3cd..f065fd26 100644 --- a/harness/tht/cli/search_cmd.py +++ b/harness/tht/cli/search_cmd.py @@ -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 diff --git a/harness/tht/jobs/dwh_pipeline.py b/harness/tht/jobs/dwh_pipeline.py index d10f4c8f..fab1e265 100644 --- a/harness/tht/jobs/dwh_pipeline.py +++ b/harness/tht/jobs/dwh_pipeline.py @@ -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: diff --git a/harness/tht/jobs/runner.py b/harness/tht/jobs/runner.py index 5c052760..89f73b46 100644 --- a/harness/tht/jobs/runner.py +++ b/harness/tht/jobs/runner.py @@ -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