mirror of
https://github.com/Abdess/retroarch_system.git
synced 2026-10-10 21:43:23 -05:00
fix: avoid shared scratch path in large-file cache
This commit is contained in:
1 parent
0a59af3279
commit
c51fc233f9
2 files changed
+140
-1
No files matched your search
+9
-1
@@ -10,6 +10,7 @@ import contextlib
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
@@ -1220,7 +1221,12 @@ def fetch_large_file(
|
||||
return cached
|
||||
|
||||
os.makedirs(dest_dir, exist_ok=True)
|
||||
tmp_path = cached + ".tmp"
|
||||
# A per-process scratch name: two runs fetching the same asset into one
|
||||
# shared path interleave their writes into a full-size, corrupt file.
|
||||
tmp_fd, tmp_path = tempfile.mkstemp(
|
||||
dir=dest_dir, prefix=os.path.basename(cached) + ".", suffix=".tmp"
|
||||
)
|
||||
os.close(tmp_fd)
|
||||
# GitHub rewrites spaces to dots in release asset names, so a file whose
|
||||
# name contains spaces is published under a dotted name.
|
||||
candidates = [name]
|
||||
@@ -1250,6 +1256,8 @@ def fetch_large_file(
|
||||
os.unlink(tmp_path)
|
||||
|
||||
if not downloaded:
|
||||
if os.path.exists(tmp_path):
|
||||
os.unlink(tmp_path)
|
||||
return None
|
||||
|
||||
if expected_sha1 or expected_md5:
|
||||
|
||||
@@ -0,0 +1,131 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Large-file cache downloads.
|
||||
|
||||
Two runs fetching the same asset used to stream into one shared scratch
|
||||
path, interleaving their writes into a full-size file with mixed content
|
||||
that then replaced the cache entry.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import io
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import unittest
|
||||
import urllib.error
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parent.parent
|
||||
sys.path.insert(0, str(REPO_ROOT / "scripts"))
|
||||
|
||||
import common # noqa: E402
|
||||
|
||||
PAYLOAD_A = b"A" * (256 * 1024)
|
||||
PAYLOAD_B = b"B" * (256 * 1024)
|
||||
|
||||
|
||||
class _SlowResponse(io.BytesIO):
|
||||
"""Serve a payload in chunks, yielding between them."""
|
||||
|
||||
def __init__(self, data: bytes, barrier: threading.Barrier | None = None):
|
||||
super().__init__(data)
|
||||
self._barrier = barrier
|
||||
self._first = True
|
||||
|
||||
def read(self, size: int = -1) -> bytes:
|
||||
chunk = super().read(min(size, 4096) if size and size > 0 else 4096)
|
||||
if self._first and self._barrier is not None:
|
||||
self._first = False
|
||||
self._barrier.wait(timeout=10)
|
||||
return chunk
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc):
|
||||
self.close()
|
||||
return False
|
||||
|
||||
|
||||
class LargeFileCacheTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.dir = tempfile.mkdtemp()
|
||||
self._urlopen = common.urllib.request.urlopen
|
||||
|
||||
def tearDown(self):
|
||||
common.urllib.request.urlopen = self._urlopen
|
||||
|
||||
def test_concurrent_fetches_do_not_mix(self):
|
||||
barrier = threading.Barrier(2, timeout=10)
|
||||
payloads = [PAYLOAD_A, PAYLOAD_B]
|
||||
index = iter(range(2))
|
||||
lock = threading.Lock()
|
||||
|
||||
def fake_urlopen(req, timeout=None):
|
||||
with lock:
|
||||
i = next(index)
|
||||
return _SlowResponse(payloads[i], barrier)
|
||||
|
||||
common.urllib.request.urlopen = fake_urlopen
|
||||
results: list[str | None] = [None, None]
|
||||
|
||||
def worker(slot: int):
|
||||
results[slot] = common.fetch_large_file("asset.bin", dest_dir=self.dir)
|
||||
|
||||
threads = [threading.Thread(target=worker, args=(i,)) for i in range(2)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=30)
|
||||
|
||||
cached = Path(self.dir) / "asset.bin"
|
||||
self.assertTrue(cached.exists())
|
||||
digest = hashlib.sha1(cached.read_bytes()).hexdigest()
|
||||
accepted = {hashlib.sha1(p).hexdigest() for p in payloads}
|
||||
# Whichever run lands last wins, but never a blend of the two
|
||||
self.assertIn(digest, accepted)
|
||||
|
||||
def test_no_scratch_file_survives_a_successful_fetch(self):
|
||||
common.urllib.request.urlopen = lambda req, timeout=None: _SlowResponse(
|
||||
PAYLOAD_A
|
||||
)
|
||||
common.fetch_large_file("asset.bin", dest_dir=self.dir)
|
||||
leftovers = [f for f in os.listdir(self.dir) if f.endswith(".tmp")]
|
||||
self.assertEqual(leftovers, [])
|
||||
|
||||
def test_no_scratch_file_survives_a_failed_fetch(self):
|
||||
def fail(req, timeout=None):
|
||||
raise urllib.error.URLError("offline")
|
||||
|
||||
common.urllib.request.urlopen = fail
|
||||
self.assertIsNone(common.fetch_large_file("asset.bin", dest_dir=self.dir))
|
||||
self.assertEqual(os.listdir(self.dir), [])
|
||||
|
||||
def test_hash_mismatch_leaves_no_scratch_file(self):
|
||||
common.urllib.request.urlopen = lambda req, timeout=None: _SlowResponse(
|
||||
PAYLOAD_A
|
||||
)
|
||||
result = common.fetch_large_file(
|
||||
"asset.bin", dest_dir=self.dir, expected_sha1="00" * 20
|
||||
)
|
||||
self.assertIsNone(result)
|
||||
self.assertEqual(os.listdir(self.dir), [])
|
||||
|
||||
def test_cached_file_is_returned_without_download(self):
|
||||
cached = Path(self.dir) / "asset.bin"
|
||||
cached.write_bytes(PAYLOAD_A)
|
||||
|
||||
def fail(req, timeout=None):
|
||||
raise AssertionError("must not download when the cache is valid")
|
||||
|
||||
common.urllib.request.urlopen = fail
|
||||
self.assertEqual(
|
||||
common.fetch_large_file("asset.bin", dest_dir=self.dir), str(cached)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in new issue
Block a user