fix(preprocess): lease DWH generation reads

This commit is contained in:
2026-07-12 05:58:12 +02:00
parent db35ddb041
commit 6bb158233f
3 changed files with 205 additions and 23 deletions
+65
View File
@@ -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
+13 -6
View File
@@ -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] = []
+127 -17
View File
@@ -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: