From 6bb158233f09d079a2739b8fa793dd5a2f19077c Mon Sep 17 00:00:00 2001 From: mptyl Date: Sun, 12 Jul 2026 05:58:12 +0200 Subject: [PATCH] fix(preprocess): lease DWH generation reads --- harness/tests/test_dwh_preprocess_job.py | 65 ++++++++++ harness/tht/cli/search_cmd.py | 19 ++- harness/tht/jobs/dwh_pipeline.py | 144 ++++++++++++++++++++--- 3 files changed, 205 insertions(+), 23 deletions(-) diff --git a/harness/tests/test_dwh_preprocess_job.py b/harness/tests/test_dwh_preprocess_job.py index aa86b1a3..4d8a76db 100644 --- a/harness/tests/test_dwh_preprocess_job.py +++ b/harness/tests/test_dwh_preprocess_job.py @@ -7,6 +7,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.dwh_pipeline import lease_dwh_snapshot from tht.jobs.locking import _lock_name @@ -228,3 +229,67 @@ def test_generation_retention_keeps_active_and_one_rollback(tmp_path): run_ids.append(report.run_id) remaining = {path.name for path in (tmp_path / ".tht-dwh" / "generations").iterdir()} assert remaining == set(run_ids[-2:]) + + +def test_reader_lease_blocks_retain_one_publisher_until_file_reads_finish(tmp_path): + import threading + import time + from types import SimpleNamespace + + def make(content, retain=1): + 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") + ], retain_generations=retain, + ) + + first = make("old").run() + cfg = SimpleNamespace(paths=SimpleNamespace( + artifacts=tmp_path / "artifacts", indexes=tmp_path / "indexes" + )) + completed = threading.Event() + with lease_dwh_snapshot(cfg) as snapshot: + thread = threading.Thread(target=lambda: (make("new").run(), completed.set())) + thread.start() + time.sleep(0.05) + assert not completed.is_set() + assert snapshot.physical.read_text() == "old" + assert snapshot.generation == first.run_id + thread.join(timeout=2) + assert completed.is_set() + assert not (tmp_path / ".tht-dwh" / "generations" / first.run_id).exists() + + +def test_cleanup_never_follows_top_level_or_child_symlinks(tmp_path): + external = tmp_path / "external" + external.mkdir() + victim = external / "victim" + victim.write_text("safe") + + generations = tmp_path / ".tht-dwh" / "generations" + generations.mkdir(parents=True) + (generations / ("a" * 32)).symlink_to(external, target_is_directory=True) + + def make(content, retain=1): + 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") + ], retain_generations=retain, + ) + + first = make("one", retain=2).run() + make("two", retain=2).run() + old = generations / first.run_id + old.chmod(0o700) + (old / "hostile-link").symlink_to(victim) + make("three").run() + assert victim.read_text() == "safe" + assert victim.stat().st_mode & 0o200 diff --git a/harness/tht/cli/search_cmd.py b/harness/tht/cli/search_cmd.py index ed6d79bf..7500e25a 100644 --- a/harness/tht/cli/search_cmd.py +++ b/harness/tht/cli/search_cmd.py @@ -21,8 +21,18 @@ DEFAULT_TOP_FALLBACK = 10 search_app = typer.Typer(help="Ricerca semantica (evidence/schema/values) nel vectorstore") +def _leased_dwh_snapshot(cfg, context: typer.Context): + from tht.jobs.dwh_pipeline import lease_dwh_snapshot + + lease = lease_dwh_snapshot(cfg) + snapshot = lease.__enter__() + context.call_on_close(lambda: lease.__exit__(None, None, None)) + return snapshot + + @search_app.command("find") def search_cmd( + ctx: typer.Context, keyword: str = typer.Argument(..., help="Termine da cercare, es. 'ablazione'."), config: Path = CONFIG_OPT, top: int | None = typer.Option( @@ -47,9 +57,7 @@ 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) + dwh_snapshot = _leased_dwh_snapshot(cfg, ctx) require_vector_cfg(cfg) from tht.search.evidence import active_searcher @@ -225,6 +233,7 @@ PACK_EXCERPT_CHARS = 400 @search_app.command("pack") def pack_cmd( + ctx: typer.Context, question: str = typer.Argument(..., help="La domanda in linguaggio naturale."), config: Path = CONFIG_OPT, session: str = typer.Option( @@ -247,9 +256,7 @@ 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) + dwh_snapshot = _leased_dwh_snapshot(cfg, ctx) require_vector_cfg(cfg) tables: list[dict] = [] diff --git a/harness/tht/jobs/dwh_pipeline.py b/harness/tht/jobs/dwh_pipeline.py index fab1e265..3581d1ec 100644 --- a/harness/tht/jobs/dwh_pipeline.py +++ b/harness/tht/jobs/dwh_pipeline.py @@ -4,6 +4,7 @@ from __future__ import annotations import hashlib import json +import fcntl import os import re import shutil @@ -36,6 +37,47 @@ class DwhArtifactSnapshot: lsh_dir: Path +class DwhSnapshotLease: + def __init__(self, cfg) -> None: + self.cfg = cfg + self._fd: int | None = None + self.snapshot: DwhArtifactSnapshot | None = None + + def __enter__(self) -> DwhArtifactSnapshot: + self._fd = _acquire_generation_lock(self.cfg.paths.artifacts.parent, exclusive=False) + try: + self.snapshot = resolve_dwh_snapshot(self.cfg) + return self.snapshot + except BaseException: + self.__exit__(None, None, None) + raise + + def __exit__(self, *_args) -> None: + if self._fd is not None: + fd, self._fd = self._fd, None + fcntl.flock(fd, fcntl.LOCK_UN) + os.close(fd) + + +def lease_dwh_snapshot(cfg) -> DwhSnapshotLease: + return DwhSnapshotLease(cfg) + + +def _acquire_generation_lock(workspace_root: Path, *, exclusive: bool) -> int: + root = workspace_root / ".tht-dwh" + DwhPreprocessPipeline._ensure_owned_dir(root) + fd = os.open(root / "generation.lock", os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW, 0o600) + try: + info = os.fstat(fd) + if not stat.S_ISREG(info.st_mode) or info.st_uid != os.getuid() or info.st_nlink != 1: + raise OSError("unsafe DWH generation lock") + fcntl.flock(fd, fcntl.LOCK_EX if exclusive else fcntl.LOCK_SH) + return fd + except BaseException: + os.close(fd) + raise + + def _digest(path: Path) -> str: return hashlib.sha256(path.read_bytes()).hexdigest() @@ -48,7 +90,7 @@ def _read_owned(path: Path, *, readonly: bool) -> bytes: 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)) + or (readonly and bool(info.st_mode & 0o222)) ): raise OSError("unsafe DWH generation file") chunks = [] @@ -62,6 +104,13 @@ def _read_owned(path: Path, *, readonly: bool) -> bytes: def validate_generation(target: Path) -> dict: manifest_path = target / GENERATION_MANIFEST try: + directory_info = target.lstat() + if ( + not stat.S_ISDIR(directory_info.st_mode) + or directory_info.st_uid != os.getuid() + or stat.S_IMODE(directory_info.st_mode) != 0o700 + ): + raise ValueError manifest = json.loads(_read_owned(manifest_path, readonly=True).decode("utf-8")) files = manifest["files"] if ( @@ -273,10 +322,15 @@ class DwhPreprocessPipeline: artifacts = self._artifacts(context) required = self._required(artifacts) seal_stage_artifacts(context, stage, required, spec) - self._publish(context.run_id, artifacts, required) - if self.after_publish is not None: - self.after_publish(context.run_id) - self._cleanup_generations() + lease_fd = _acquire_generation_lock(self.workspace_root, exclusive=True) + try: + self._publish(context.run_id, artifacts, required) + if self.after_publish is not None: + self.after_publish(context.run_id) + self._cleanup_generations() + finally: + fcntl.flock(lease_fd, fcntl.LOCK_UN) + os.close(lease_fd) return StageArtifacts(required) def _publish(self, generation: str, artifacts: Path, required: tuple[str, ...]) -> None: @@ -388,19 +442,75 @@ class DwhPreprocessPipeline: 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: + generations = [] + for path in root.iterdir(): + try: + info = path.lstat() + except OSError: continue - path.chmod(0o700) - for child in path.iterdir(): - child.chmod(0o600) - shutil.rmtree(path) - self._fsync(root) + if ( + _RUN_ID.fullmatch(path.name) + and stat.S_ISDIR(info.st_mode) + and info.st_uid == os.getuid() + ): + generations.append((info.st_mtime_ns, path.name)) + generations.sort() + keep_recent = {name for _, name in generations[-self.retain_generations:]} + root_fd = os.open(root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW) + try: + for _, name in generations: + if name in protected | keep_recent: + continue + self._safe_delete_generation(root_fd, name) + os.fsync(root_fd) + finally: + os.close(root_fd) + + @staticmethod + def _safe_delete_generation(root_fd: int, name: str) -> None: + try: + generation_fd = os.open( + name, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW, dir_fd=root_fd + ) + except OSError: + return + try: + info = os.fstat(generation_fd) + if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid(): + return + entries = os.listdir(generation_fd) + opened = [] + try: + for child in entries: + try: + fd = os.open( + child, os.O_RDONLY | os.O_NOFOLLOW, dir_fd=generation_fd + ) + except OSError: + return + child_info = os.fstat(fd) + if ( + not stat.S_ISREG(child_info.st_mode) + or child_info.st_uid != os.getuid() + or child_info.st_nlink != 1 + ): + os.close(fd) + return + opened.append((child, fd)) + for child, fd in opened: + os.fchmod(fd, 0o600) + os.close(fd) + os.unlink(child, dir_fd=generation_fd) + opened.clear() + finally: + for _, fd in opened: + os.close(fd) + finally: + os.close(generation_fd) + try: + os.rmdir(name, dir_fd=root_fd) + except OSError: + return @staticmethod def _ensure_owned_dir(path: Path) -> None: