fix(preprocess): claim DWH roots atomically

This commit is contained in:
2026-07-12 06:51:46 +02:00
parent 8110793f61
commit c4dd6c8900
2 changed files with 199 additions and 10 deletions
+75 -7
View File
@@ -104,17 +104,87 @@ def test_shared_root_mismatch_fails_without_deadlock_while_owner_reader_is_activ
target=lambda: (errors.append(_capture_error(contender.run)), finished.set())
)
thread.start()
assert finished.wait(2)
assert not finished.wait(0.1)
thread.join(2)
assert finished.is_set()
assert "different workspace configuration" in str(errors[0])
def test_concurrent_brand_new_shared_root_has_one_atomic_owner_and_loser_never_builds(tmp_path):
import threading
calls = {"alpha": 0, "beta": 0}
results = []
barrier = threading.Barrier(2)
def run(workspace):
def introspect(output):
calls[workspace] += 1
output.write_text("catalog")
def build(physical, output):
calls[workspace] += 1
_write_lsh([], physical, output)
candidate = DwhPreprocessPipeline(
workspace_id=workspace, workspace_root=tmp_path,
config_fingerprint=FP, input_fingerprint=FP,
introspect=introspect, build_lsh=build,
)
barrier.wait()
results.append((workspace, _capture_error(candidate.run)))
threads = [threading.Thread(target=run, args=(name,)) for name in ("alpha", "beta")]
for thread in threads:
thread.start()
for thread in threads:
thread.join(5)
assert all(not thread.is_alive() for thread in threads)
winner = next(name for name, result in results if not isinstance(result, Exception))
loser = next(name for name, result in results if isinstance(result, Exception))
assert calls[winner] == 2
assert calls[loser] == 0
def test_missing_active_with_generations_and_symlink_owner_marker_fail_closed(tmp_path):
import pytest
calls = []
owner = 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: _write_lsh([], physical, output),
)
owner.run()
(tmp_path / ".tht-dwh" / "ACTIVE").unlink()
owner.introspect = lambda output: calls.append("called")
with pytest.raises(Exception, match="without a consistent ACTIVE"):
owner.run()
assert calls == []
other_root = tmp_path / "other"
marker_root = other_root / ".tht-dwh"
marker_root.mkdir(parents=True)
external = tmp_path / "external-owner"
external.write_text("foreign")
(marker_root / "OWNER.json").symlink_to(external)
contender = DwhPreprocessPipeline(
workspace_id="demo", workspace_root=other_root,
config_fingerprint=FP, input_fingerprint=FP,
introspect=lambda output: calls.append("symlink-called"),
build_lsh=lambda physical, output: None,
)
with pytest.raises(Exception, match="ownership marker"):
contender.run()
assert calls == []
def _capture_error(operation):
try:
operation()
return operation()
except Exception as error:
return error
raise AssertionError("operation unexpectedly succeeded")
def _write_lsh(calls, physical: Path, output: Path):
@@ -430,10 +500,6 @@ def test_cleanup_never_follows_top_level_or_child_symlinks(tmp_path):
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,
@@ -446,6 +512,8 @@ def test_cleanup_never_follows_top_level_or_child_symlinks(tmp_path):
)
first = make("one", retain=2).run()
generations = tmp_path / ".tht-dwh" / "generations"
(generations / ("a" * 32)).symlink_to(external, target_is_directory=True)
make("two", retain=2).run()
old = generations / first.run_id
old.chmod(0o700)
+124 -3
View File
@@ -28,6 +28,7 @@ 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"
OWNER_MARKER = "OWNER.json"
def config_dwh_binding(cfg) -> dict[str, str]:
@@ -42,6 +43,96 @@ def config_dwh_binding(cfg) -> dict[str, str]:
}
def _binding_digest(binding: dict[str, str]) -> str:
payload = json.dumps(binding, sort_keys=True, separators=(",", ":"))
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
def _read_root_binding(workspace_root: Path) -> dict[str, str]:
marker = workspace_root / ".tht-dwh" / OWNER_MARKER
try:
payload = json.loads(_read_owned(marker, readonly=True).decode("utf-8"))
binding = payload["binding"]
if (
payload.get("schema_version") != 1
or not isinstance(binding, dict)
or set(binding) != {"workspace_id", "config_fingerprint", "input_fingerprint"}
or payload.get("binding_sha256") != _binding_digest(binding)
):
raise ValueError
return binding
except (OSError, KeyError, TypeError, ValueError, UnicodeDecodeError) as error:
raise CorruptCheckpointError("DWH workspace ownership marker is missing or invalid") from error
def _validate_root_binding(workspace_root: Path, expected: dict[str, str]) -> None:
if _read_root_binding(workspace_root) != expected:
raise CorruptCheckpointError("DWH artifacts belong to a different workspace configuration")
root = workspace_root / ".tht-dwh"
generations = root / "generations"
active_exists = (root / "ACTIVE").exists()
try:
generations_fd = os.open(
generations, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW
)
except FileNotFoundError:
generation_entries = []
except OSError as error:
raise CorruptCheckpointError("DWH generations directory is invalid") from error
else:
try:
generation_entries = os.listdir(generations_fd)
finally:
os.close(generations_fd)
if generation_entries and not active_exists:
raise CorruptCheckpointError("DWH generations exist without a consistent ACTIVE pointer")
def _claim_or_validate_root_binding(workspace_root: Path, binding: dict[str, str]) -> None:
root = workspace_root / ".tht-dwh"
marker = root / OWNER_MARKER
if marker.exists() or marker.is_symlink():
_validate_root_binding(workspace_root, binding)
return
generations = root / "generations"
generations_nonempty = False
try:
generations_fd = os.open(
generations, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW
)
except FileNotFoundError:
pass
except OSError as error:
raise CorruptCheckpointError("DWH generations directory is invalid") from error
else:
try:
generations_nonempty = bool(os.listdir(generations_fd))
finally:
os.close(generations_fd)
if (root / "ACTIVE").exists() or generations_nonempty:
raise CorruptCheckpointError(
"DWH artifacts are unbound; migrate them explicitly or use an empty root"
)
payload = {
"schema_version": 1,
"binding": binding,
"binding_sha256": _binding_digest(binding),
}
temporary = root / f".{OWNER_MARKER}.{uuid.uuid4().hex}.tmp"
fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW, 0o600)
try:
with os.fdopen(fd, "w", encoding="utf-8") as stream:
stream.write(json.dumps(payload, sort_keys=True, separators=(",", ":")) + "\n")
stream.flush()
os.fsync(stream.fileno())
temporary.chmod(0o400)
os.replace(temporary, marker)
DwhPreprocessPipeline._fsync(root)
except BaseException:
temporary.unlink(missing_ok=True)
raise
@dataclass(frozen=True)
class DwhArtifactSnapshot:
generation: str | None
@@ -56,9 +147,16 @@ class DwhSnapshotLease:
self.snapshot: DwhArtifactSnapshot | None = None
def __enter__(self) -> DwhArtifactSnapshot:
binding = config_dwh_binding(self.cfg)
claim_fd = _acquire_generation_lock(self.cfg.paths.artifacts.parent, exclusive=True)
try:
_claim_or_validate_root_binding(self.cfg.paths.artifacts.parent, binding)
finally:
fcntl.flock(claim_fd, fcntl.LOCK_UN)
os.close(claim_fd)
self._fd = _acquire_generation_lock(self.cfg.paths.artifacts.parent, exclusive=False)
try:
self.snapshot = resolve_dwh_snapshot(self.cfg)
self.snapshot = _resolve_dwh_snapshot_locked(self.cfg, binding)
return self.snapshot
except BaseException:
self.__exit__(None, None, None)
@@ -190,12 +288,31 @@ def validate_generation(
def resolve_dwh_snapshot(cfg) -> DwhArtifactSnapshot:
binding = config_dwh_binding(cfg)
claim_fd = _acquire_generation_lock(cfg.paths.artifacts.parent, exclusive=True)
try:
_claim_or_validate_root_binding(cfg.paths.artifacts.parent, binding)
finally:
fcntl.flock(claim_fd, fcntl.LOCK_UN)
os.close(claim_fd)
lease_fd = _acquire_generation_lock(cfg.paths.artifacts.parent, exclusive=False)
try:
return _resolve_dwh_snapshot_locked(cfg, binding)
finally:
fcntl.flock(lease_fd, fcntl.LOCK_UN)
os.close(lease_fd)
def _resolve_dwh_snapshot_locked(
cfg, binding: dict[str, str],
) -> DwhArtifactSnapshot:
_validate_root_binding(cfg.paths.artifacts.parent, binding)
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, config_dwh_binding(cfg))
validate_generation(target, binding)
return DwhArtifactSnapshot(target.name, target / "physical.yaml", target)
@@ -263,6 +380,7 @@ class DwhPreprocessPipeline:
}
def _assert_active_binding(self) -> None:
_validate_root_binding(self.workspace_root, self.binding)
active = active_generation_dir(self.workspace_root)
if active is not None:
validate_generation(active, self.binding)
@@ -271,8 +389,9 @@ class DwhPreprocessPipeline:
self, steps: tuple[str, ...] = DWH_STAGE_IDS, *, resume_run_id: str | None = None
) -> JobReport:
self._validate_steps(steps)
lease_fd = _acquire_generation_lock(self.workspace_root, exclusive=False)
lease_fd = _acquire_generation_lock(self.workspace_root, exclusive=True)
try:
_claim_or_validate_root_binding(self.workspace_root, self.binding)
self._assert_active_binding()
finally:
fcntl.flock(lease_fd, fcntl.LOCK_UN)
@@ -399,6 +518,7 @@ class DwhPreprocessPipeline:
seal_stage_artifacts(context, stage, required, spec)
lease_fd = _acquire_generation_lock(self.workspace_root, exclusive=True)
try:
_validate_root_binding(self.workspace_root, self.binding)
self._assert_active_binding()
self._publish(context.run_id, artifacts, required)
if self.after_publish is not None:
@@ -509,6 +629,7 @@ class DwhPreprocessPipeline:
root = self.workspace_root / ".tht-dwh" / "generations"
if not root.exists():
return
_validate_root_binding(self.workspace_root, self.binding)
active = active_generation_dir(self.workspace_root)
if active is not None:
validate_generation(active, self.binding)