mirror of
https://github.com/Abdess/retroarch_system.git
synced 2026-10-10 13:33:24 -05:00
fix: refresh data dirs under a lock, by rename
This commit is contained in:
1 parent
9e4b321886
commit
7da259ec6a
2 files changed
+217
-39
No files matched your search
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in new issue
Block a user