diff --git a/scripts/common.py b/scripts/common.py index c6a287cb..9db5d4f3 100644 --- a/scripts/common.py +++ b/scripts/common.py @@ -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: diff --git a/tests/test_large_file_cache.py b/tests/test_large_file_cache.py new file mode 100644 index 00000000..b47e2182 --- /dev/null +++ b/tests/test_large_file_cache.py @@ -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()