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
+119
View File
@@ -0,0 +1,119 @@
"""Data directory refreshes that other sessions can run at the same time.
Two refreshes of one key shared a fixed `<key>.previous` name and a
non-atomic existence test, so the second tree landed inside the first as
`extract/`. Staging in /tmp turned the swap into a 228 MB copy, and the
version file was truncated in place, where a concurrent reader died on an
empty JSON document.
"""
from __future__ import annotations
import io
import json
import sys
import tempfile
import threading
import unittest
import zipfile
from pathlib import Path
from unittest import mock
REPO_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(REPO_ROOT / "scripts"))
import refresh_data_dirs as rdd # noqa: E402
def _zip_bytes(files: dict[str, bytes]) -> bytes:
buffer = io.BytesIO()
with zipfile.ZipFile(buffer, "w") as zf:
for name, data in files.items():
zf.writestr(name, data)
return buffer.getvalue()
class _Response(io.BytesIO):
def __enter__(self):
return self
def __exit__(self, *exc):
return False
class RefreshConcurrency(unittest.TestCase):
def setUp(self):
(REPO_ROOT / "tmp").mkdir(exist_ok=True)
self.tmp = tempfile.TemporaryDirectory(dir=REPO_ROOT / "tmp")
self.addCleanup(self.tmp.cleanup)
self.root = Path(self.tmp.name)
self.cache = self.root / "data" / "pak"
self.versions = str(self.root / "data" / ".versions.json")
def test_tree_is_swapped_in_by_rename(self):
payload = _zip_bytes({"a.bin": b"new"})
self.cache.mkdir(parents=True)
(self.cache / "a.bin").write_bytes(b"old")
with mock.patch.object(
rdd.urllib.request, "urlopen", lambda *a, **k: _Response(payload)
), mock.patch.object(
rdd.shutil, "move", side_effect=AssertionError("copy across filesystems")
):
count = rdd._download_and_extract_zip("https://x/pak.zip", str(self.cache))
self.assertEqual(count, 1)
self.assertEqual((self.cache / "a.bin").read_bytes(), b"new")
leftovers = [p.name for p in self.cache.parent.iterdir() if p.name != "pak"]
self.assertEqual(leftovers, [])
def test_refresh_holds_the_directory_lock(self):
import fcntl
payload = _zip_bytes({"a.bin": b"x"})
lock_path = self.cache.with_name(".pak.lock")
seen: list[bool] = []
def fake_urlopen(*a, **k):
with open(lock_path, "w") as handle:
try:
fcntl.flock(handle, fcntl.LOCK_EX | fcntl.LOCK_NB)
except BlockingIOError:
seen.append(True)
else:
fcntl.flock(handle, fcntl.LOCK_UN)
seen.append(False)
return _Response(payload)
entry = {"source_type": "zip", "source_url": "https://x/pak.zip",
"local_cache": str(self.cache)}
with mock.patch.object(rdd.urllib.request, "urlopen", fake_urlopen), \
mock.patch.object(rdd, "_get_remote_etag", return_value="e1"):
self.assertTrue(
rdd.refresh_entry("pak", entry, force=True, versions_path=self.versions)
)
self.assertEqual(seen, [True])
def test_version_file_is_never_left_truncated(self):
Path(self.versions).parent.mkdir(parents=True)
Path(self.versions).write_text(json.dumps({"keep": {"sha": "1"}}))
with mock.patch.object(rdd.json, "dump", side_effect=OSError("disk full")):
with self.assertRaises(OSError):
rdd._save_versions({"keep": {"sha": "2"}}, self.versions)
self.assertEqual(json.loads(Path(self.versions).read_text()), {"keep": {"sha": "1"}})
def test_concurrent_writers_keep_each_other(self):
keys = [f"k{i}" for i in range(16)]
threads = [
threading.Thread(
target=rdd._record_version, args=(k, {"sha": k}, self.versions)
)
for k in keys
]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
self.assertEqual(sorted(rdd._load_versions(self.versions)), sorted(keys))
if __name__ == "__main__":
unittest.main()