From 58cb61a08fae40141821163f5a48c85665c82cd7 Mon Sep 17 00:00:00 2001 From: Abdessamad Derraz <3028866+Abdess@users.noreply.github.com> Date: Tue, 6 Oct 2026 01:22:56 +0200 Subject: [PATCH] fix: refresh data dirs under a lock, by rename --- scripts/refresh_data_dirs.py | 137 +++++++++++++++++++++++--------- tests/test_refresh_data_dirs.py | 119 +++++++++++++++++++++++++++ 2 files changed, 217 insertions(+), 39 deletions(-) create mode 100644 tests/test_refresh_data_dirs.py diff --git a/scripts/refresh_data_dirs.py b/scripts/refresh_data_dirs.py index 3775d812..3eaaaadd 100644 --- a/scripts/refresh_data_dirs.py +++ b/scripts/refresh_data_dirs.py @@ -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 diff --git a/tests/test_refresh_data_dirs.py b/tests/test_refresh_data_dirs.py new file mode 100644 index 00000000..15ce8648 --- /dev/null +++ b/tests/test_refresh_data_dirs.py @@ -0,0 +1,119 @@ +"""Data directory refreshes that other sessions can run at the same time. + +Two refreshes of one key shared a fixed `.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()