fix: refresh data dirs under a lock, by rename

This commit is contained in:
Abdessamad Derraz committed 2026-10-06 01:22:56 +02:00
1 parent 485e5bc2b3
commit 58cb61a08f
2 files changed
+217 -39

No files matched your search

+98 -39
View File
@@ -13,6 +13,7 @@ Usage:
from __future__ import annotations
import argparse
import contextlib
import json
import logging
import os
@@ -50,6 +51,28 @@ def load_registry(registry_path: str = DEFAULT_REGISTRY) -> dict[str, dict]:
return data.get("data_directories", {})
@contextlib.contextmanager
def _file_lock(lock_path: Path):
"""Hold an exclusive lock on lock_path, waiting for it if taken.
Several sessions refresh the same data directories: without it, two
swaps of one tree interleave and the second lands inside the first.
On platforms without flock the lock is a no-op.
"""
try:
import fcntl
except ImportError:
yield
return
lock_path.parent.mkdir(parents=True, exist_ok=True)
with open(lock_path, "w") as handle:
fcntl.flock(handle, fcntl.LOCK_EX)
try:
yield
finally:
fcntl.flock(handle, fcntl.LOCK_UN)
def _load_versions(versions_path: str = VERSIONS_FILE) -> dict[str, dict]:
path = Path(versions_path)
if not path.exists():
@@ -61,11 +84,65 @@ def _load_versions(versions_path: str = VERSIONS_FILE) -> dict[str, dict]:
def _save_versions(
versions: dict[str, dict], versions_path: str = VERSIONS_FILE
) -> None:
"""Write the version file whole or not at all.
Truncating in place let a concurrent reader load an empty file and die
on JSONDecodeError. The scratch file sits beside the target so the
rename stays on one filesystem.
"""
path = Path(versions_path)
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w") as f:
json.dump(versions, f, indent=2, sort_keys=True)
f.write("\n")
handle, scratch = tempfile.mkstemp(
dir=path.parent, prefix=f".{path.name}.", suffix=".tmp"
)
try:
with os.fdopen(handle, "w") as f:
json.dump(versions, f, indent=2, sort_keys=True)
f.write("\n")
os.replace(scratch, path)
except BaseException:
with contextlib.suppress(OSError):
os.unlink(scratch)
raise
def _record_version(key: str, record: dict, versions_path: str) -> None:
"""Set one key under the file lock, so two writers keep both updates."""
path = Path(versions_path)
with _file_lock(path.with_name(f".{path.name}.lock")):
versions = _load_versions(versions_path)
versions[key] = record
_save_versions(versions, versions_path)
def _staging_dir(cache_dir: Path) -> tempfile.TemporaryDirectory:
"""Scratch space beside the cache, on its filesystem.
/tmp is a 4 GB tmpfs here: staging there filled it and turned the
promotion into a copy, during which the cache held a partial tree.
"""
cache_dir.parent.mkdir(parents=True, exist_ok=True)
return tempfile.TemporaryDirectory(
dir=cache_dir.parent, prefix=f".{cache_dir.name}-"
)
def _promote(extract_dir: Path, cache_dir: Path, scratch: Path) -> None:
"""Swap the new tree in by renames alone.
The old tree steps aside into the run's own scratch directory, never a
fixed sibling name another run could be using, and comes back if the
new one cannot be moved in.
"""
previous = scratch / "previous"
if cache_dir.exists():
os.replace(cache_dir, previous)
try:
os.replace(extract_dir, cache_dir)
except BaseException:
if previous.exists() and not cache_dir.exists():
os.replace(previous, cache_dir)
raise
def _api_request(url: str) -> dict:
@@ -149,7 +226,7 @@ def _download_and_extract(
exclude = exclude or []
cache_dir = Path(local_cache)
with tempfile.TemporaryDirectory() as tmpdir:
with _staging_dir(cache_dir) as tmpdir:
tarball_path = Path(tmpdir) / "archive.tar.gz"
log.info("downloading %s", source_url)
@@ -206,22 +283,7 @@ def _download_and_extract(
shutil.copyfileobj(src, dst)
file_count += 1
# atomic swap: rename old before moving new into place
cache_dir.parent.mkdir(parents=True, exist_ok=True)
old_cache = cache_dir.with_suffix(".old")
if cache_dir.exists():
if old_cache.exists():
shutil.rmtree(old_cache)
cache_dir.rename(old_cache)
try:
shutil.move(str(extract_dir), str(cache_dir))
except OSError:
# Restore old cache on failure
if old_cache.exists() and not cache_dir.exists():
old_cache.rename(cache_dir)
raise
if old_cache.exists():
shutil.rmtree(old_cache)
_promote(extract_dir, cache_dir, Path(tmpdir))
return file_count
@@ -241,7 +303,7 @@ def _download_and_extract_zip(
exclude = exclude or []
cache_dir = Path(local_cache)
with tempfile.TemporaryDirectory() as tmpdir:
with _staging_dir(cache_dir) as tmpdir:
zip_path = Path(tmpdir) / "archive.zip"
log.info("downloading %s", source_url)
@@ -285,21 +347,7 @@ def _download_and_extract_zip(
# The old tree is stepped aside rather than deleted: removing it
# first and then failing to move the new one in left the cache with
# nothing at all, and the next run reads that as "never fetched".
cache_dir.parent.mkdir(parents=True, exist_ok=True)
previous = None
if cache_dir.exists():
previous = cache_dir.with_name(cache_dir.name + ".previous")
if previous.exists():
shutil.rmtree(previous)
os.replace(cache_dir, previous)
try:
shutil.move(str(extract_dir), str(cache_dir))
except BaseException:
if previous is not None:
os.replace(previous, cache_dir)
raise
if previous is not None:
shutil.rmtree(previous, ignore_errors=True)
_promote(extract_dir, cache_dir, Path(tmpdir))
return file_count
@@ -327,7 +375,20 @@ def refresh_entry(
"""Refresh a single data directory entry.
Returns True if the entry was refreshed (or would be in dry-run mode).
The decision and the swap run under one lock per directory: a run that
waited for another session re-reads the version it recorded and finds
nothing left to do.
"""
if dry_run:
return _refresh_entry(key, entry, force, dry_run, versions_path)
cache_dir = Path(entry["local_cache"])
with _file_lock(cache_dir.with_name(f".{cache_dir.name}.lock")):
return _refresh_entry(key, entry, force, dry_run, versions_path)
def _refresh_entry(
key: str, entry: dict, force: bool, dry_run: bool, versions_path: str
) -> bool | None:
source_type = entry.get("source_type", "tarball")
version = entry.get("version", "master")
source_url = entry["source_url"].format(version=version)
@@ -389,9 +450,7 @@ def refresh_entry(
remote_tag = _get_remote_etag(source_url)
else:
remote_tag = get_remote_sha(entry["source_url"], version)
versions = _load_versions(versions_path)
versions[key] = {"sha": remote_tag or "", "version": version}
_save_versions(versions, versions_path)
_record_version(key, {"sha": remote_tag or "", "version": version}, versions_path)
log.info("[%s] refreshed: %d files extracted to %s", key, file_count, local_cache)
return True