From 3306ee946d4050cb7a530386e1ef948a58b4e31b Mon Sep 17 00:00:00 2001 From: Codex Date: Thu, 1 Oct 2026 06:06:45 +0800 Subject: [PATCH 1/2] fix(master): publish downloads safely across concurrent writers --- src/nnnotes/master.py | 12 +++---- tests/test_master_atomic.py | 72 +++++++++++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+), 8 deletions(-) create mode 100644 tests/test_master_atomic.py diff --git a/src/nnnotes/master.py b/src/nnnotes/master.py index bf412e1..cc65f56 100644 --- a/src/nnnotes/master.py +++ b/src/nnnotes/master.py @@ -17,7 +17,6 @@ import gzip import hashlib import json -import os import re import time import urllib.error @@ -29,6 +28,7 @@ import numpy as np from . import jsonio +from .store import write_file PREFIX = 64 # bytes before the ciphertext BLOCK = 32 # 256-bit block @@ -246,8 +246,7 @@ def download(cdn: str, version: str, out_dir: Path, workers: int = 16, *, get=No seen.add(name) else: # Preserve the original partial-download receipt even if a worker raises. - from .cache import write_atomic - write_atomic(out_dir / "MasterManifest.json", manifest_raw) + write_file(out_dir / "MasterManifest.json", manifest_raw) def one(f: dict): name, sha = f["name"], (f.get("hash") or "").lower() @@ -262,16 +261,13 @@ def one(f: dict): return {"file": name, "error": "size or hash missing/different from the manifest"} if sha and hashlib.sha256(data).hexdigest() != sha: return {"file": name, "error": "sha256 differs from the manifest"} - tmp = dst.with_name(f"{name}.{os.getpid()}.part") - tmp.write_bytes(data) - os.replace(tmp, dst) + write_file(dst, data) return "downloaded" with ThreadPoolExecutor(max(1, workers)) as ex: results = list(ex.map(one, files)) if not strict or not any(isinstance(r, dict) for r in results): - from .cache import write_atomic - write_atomic(out_dir / "MasterManifest.json", manifest_raw) + write_file(out_dir / "MasterManifest.json", manifest_raw) return {"version": manifest.get("version", version), "files": len(files), "downloaded": results.count("downloaded"), "kept": results.count("kept"), "failed": [r for r in results if isinstance(r, dict)], "out": str(out_dir)} diff --git a/tests/test_master_atomic.py b/tests/test_master_atomic.py new file mode 100644 index 0000000..854d8a0 --- /dev/null +++ b/tests/test_master_atomic.py @@ -0,0 +1,72 @@ +"""Master downloads publish verified bytes using independent temporary files.""" +import threading +from concurrent.futures import ThreadPoolExecutor + +import pytest + +import synth +from nnnotes import cache, master + + +def test_concurrent_downloads_use_independent_temporary_files(tmp_path, monkeypatch): + cdn = synth.serve_master_version(tmp_path / "cdn", "v", {"MasterA.bin": b"table"}) + out = tmp_path / "out" + barrier = threading.Barrier(2) + replace = cache.os.replace + entered = set() + lock = threading.Lock() + + def synchronized_replace(src, dst): + ident = threading.get_ident() + with lock: + synchronize = dst.name == "MasterA.bin" and ident not in entered + if synchronize: + entered.add(ident) + if synchronize: + barrier.wait(timeout=5) + return replace(src, dst) + + monkeypatch.setattr(cache.os, "replace", synchronized_replace) + with ThreadPoolExecutor(2) as pool: + results = list(pool.map(lambda _: master.download(cdn, "v", out, workers=1), range(2))) + assert [r["downloaded"] for r in results] == [1, 1] + assert (out / "MasterA.bin").read_bytes() == b"table" + assert sorted(p.name for p in out.iterdir()) == ["MasterA.bin", "MasterManifest.json"] + + +def test_failed_download_write_preserves_destination_and_removes_temp(tmp_path, monkeypatch): + cdn = synth.serve_master_version(tmp_path / "cdn", "v", {"MasterA.bin": b"new table"}) + out = tmp_path / "out" + out.mkdir() + (out / "MasterA.bin").write_bytes(b"previous table") + replace = cache.os.replace + + def fail_table(src, dst): + if dst.name == "MasterA.bin": + raise OSError("simulated write failure") + return replace(src, dst) + + monkeypatch.setattr(cache.os, "replace", fail_table) + with pytest.raises(OSError, match="simulated write failure"): + master.download(cdn, "v", out, workers=1) + assert (out / "MasterA.bin").read_bytes() == b"previous table" + assert sorted(p.name for p in out.iterdir()) == ["MasterA.bin", "MasterManifest.json"] + +def test_transient_windows_rename_failure_is_retried(tmp_path, monkeypatch): + cdn = synth.serve_master_version(tmp_path / "cdn", "v", {"MasterA.bin": b"table"}) + replace = cache.os.replace + attempts = [] + + def temporarily_busy(src, dst): + if dst.name == "MasterA.bin": + attempts.append(src) + if len(attempts) == 1: + raise PermissionError("temporarily busy") + return replace(src, dst) + + monkeypatch.setattr(cache.os, "replace", temporarily_busy) + out = tmp_path / "out" + assert master.download(cdn, "v", out, workers=1)["downloaded"] == 1 + assert len(attempts) == 2 + assert (out / "MasterA.bin").read_bytes() == b"table" + assert not list(out.glob("*.part")) From 2f6177962122a81d9556191f4db0dcf3fd00a8a8 Mon Sep 17 00:00:00 2001 From: Codex Date: Thu, 1 Oct 2026 06:06:45 +0800 Subject: [PATCH 2/2] fix(catalog): reject corrupt metadata and bundle cache entries --- src/nnnotes/addressables.py | 84 ++++++++++++---------- src/nnnotes/catalog.py | 34 ++++++--- tests/test_catalog_corruption.py | 119 +++++++++++++++++++++++++++++++ 3 files changed, 189 insertions(+), 48 deletions(-) create mode 100644 tests/test_catalog_corruption.py diff --git a/src/nnnotes/addressables.py b/src/nnnotes/addressables.py index a04fc52..ac00501 100644 --- a/src/nnnotes/addressables.py +++ b/src/nnnotes/addressables.py @@ -88,37 +88,26 @@ def remote_path(internal_id: str) -> str | None: def parse(data: bytes) -> list[dict]: """Every location of a binary catalog: {offset, primary_key, internal_id, dependencies (location offsets)}.""" data = catalog_bytes(data) - def u32(offset): - return struct.unpack_from(" None: + if offset < 0 or size < 0 or offset > len(self.data) - size: + raise ValueError(f"catalog data at {offset}: past the end of the catalog") + + def unpack(self, fmt: str, offset: int) -> tuple: + self.check(offset, struct.calcsize(fmt)) + return struct.unpack_from(fmt, self.data, offset) + def u32(self, offset: int) -> int: - return struct.unpack_from(" list[int]: """A u32 array (its byte length is the u32 before it); [] for null.""" @@ -159,13 +156,15 @@ def array(self, offset: int) -> list[int]: n = self.u32(offset - 4) if n % 4: raise ValueError(f"catalog array at {offset}: byte length {n} is not a multiple of 4") - return list(struct.unpack_from(f"<{n // 4}I", self.data, offset)) + return list(self.unpack(f"<{n // 4}I", offset)) def plain(self, offset: int) -> str: s = self._plain.get(offset) if s is None: pos = offset & OFFSET_MASK - raw = self.data[pos:pos + self.u32(pos - 4)] + size = self.u32(pos - 4) + self.check(pos, size) + raw = self.data[pos:pos + size] s = self._plain[offset] = raw.decode("utf-16-le" if offset & UNICODE else "ascii") return s @@ -183,7 +182,7 @@ def string(self, offset: int, sep: str) -> str | None: if pos in seen: raise ValueError(f"catalog string at {offset & OFFSET_MASK}: the part chain loops") seen.add(pos) - part, link = struct.unpack_from(" tuple[str | None, str | None] | None: return None t = self._types.get(offset) if t is None: - assembly, cls = struct.unpack_from(" dict | None: TypeSerializer.Data.""" if offset == NONE: return None - ident, typ, data = struct.unpack_from(" tuple[str | None, object]: """A key object (ObjectTypeData {type, object}) -> (type name, value). System.String keys are an ObjectToStringRemap {u32 string, u16 separator}; System.Int32 keys a 4-byte integer; other types None.""" - typ, obj = struct.unpack_from(" dict | None: @@ -225,7 +224,7 @@ def extra(self, offset: int) -> dict | None: AssetBundleRequestOptions, the decoded options; None for null.""" if offset == NONE: return None - typ, obj = struct.unpack_from(" dict: """AssetBundleRequestOptionsSerializationAdapter.SerializedData {u32 hash, u32 bundleName, u32 crc, u32 bundleSize, u32 common}: hash a Hash128 (16 raw bytes, hex in stored order), bundleName read with '_', common a SerializedData.Common {i16 timeout, u8 redirectLimit, u8 retryCount, i32 flags}.""" - hash_id, name_id, crc, size, common = struct.unpack_from("<5I", self.data, offset) + hash_id, name_id, crc, size, common = self.unpack("<5I", offset) + if hash_id != NONE: + self.check(hash_id, 16) out = {"hash": None if hash_id == NONE else self.data[hash_id:hash_id + 16].hex(), "bundleName": self.string(name_id, "_"), "crc": crc, "bundleSize": size} if common == NONE: out.update(timeout=None, redirectLimit=None, retryCount=None, flags=None) return out - timeout, redirects, retries, flags = struct.unpack_from(" list[dict]: n = buf.u32(table - 4) if n % 8: raise ValueError(f"catalog key table: byte length {n} is not a multiple of 8") + buf.check(table, n) for pos in range(table, table + n, 8): key_obj, locations = struct.unpack_from(" list[dict]: buf = _Buffer(data) offsets: set[int] = set() table = header["keysOffset"] - for pos in range(table, table + buf.u32(table - 4), 8): + n = buf.u32(table - 4) + if n % 8: + raise ValueError(f"catalog key table: byte length {n} is not a multiple of 8") + buf.check(table, n) + for pos in range(table, table + n, 8): offsets.update(buf.array(buf.u32(pos + 4))) out = [] for pos in sorted(offsets): diff --git a/src/nnnotes/catalog.py b/src/nnnotes/catalog.py index 69877be..fb4afb0 100644 --- a/src/nnnotes/catalog.py +++ b/src/nnnotes/catalog.py @@ -16,7 +16,7 @@ from pathlib import Path from .apkset import ApkSet -from .addressables import REMOTE_PREFIX, BundleKey, decrypt, parse, parse_locations, remote_path +from .addressables import UNITYFS, REMOTE_PREFIX, BundleKey, decrypt, parse, parse_locations, remote_path from .cache import write_atomic as _write_atomic from .config import ConfigError, apk_missing, check_catalog_version @@ -59,11 +59,26 @@ def file_name(internal_id: str) -> str: def _unityfs(data: bytes, name: str, key: BundleKey | None) -> bytes: - if data[:7] == b"UnityFS": + if data.startswith(UNITYFS): return data if key is None: raise RuntimeError(f"bundle {name} is encrypted and no bundle key was given") - return decrypt(data, name, key) + decoded = decrypt(data, name, key) + if not decoded.startswith(UNITYFS): + raise ValueError(f"bundle {name} is not UnityFS after decryption (wrong key or damaged file)") + return decoded + + +def _cached_bundle(path: Path) -> Path | None: + """A decrypted bundle cache hit requires the complete UnityFS signature.""" + if path.is_file(): + try: + with path.open("rb") as stream: + if stream.read(len(UNITYFS)) == UNITYFS: + return path + except FileNotFoundError: + pass + return None class Catalog: @@ -281,7 +296,7 @@ def resolve(self, key: str) -> list[Bundle]: def cached(self, b: Bundle) -> Path | None: """The file fetch(b) returns when the bundle is in the cache already, else None.""" dst = (self.cache_dir if b.remote else self.local_cache_dir()) / "bundles" / b.name - return dst if dst.is_file() and dst.stat().st_size > 0 else None + return _cached_bundle(dst) def cached_raw(self, e: dict) -> Path | None: """The file fetch_raw(e) returns when it is in the cache already, else None.""" @@ -293,8 +308,9 @@ def fetch(self, b: Bundle) -> Path: """Local path to the decrypted bundle (CDN download or APK read).""" dst = (self.cache_dir if b.remote else self.local_cache_dir()) / "bundles" / b.name dst.parent.mkdir(parents=True, exist_ok=True) - if dst.exists() and dst.stat().st_size > 0: - return dst + hit = _cached_bundle(dst) + if hit is not None: + return hit if b.remote: url = self._url(b.internal_id) key = self._setting("bundle_key", b.name) # before the download: a missing key fails first @@ -306,7 +322,7 @@ def fetch(self, b: Bundle) -> Path: with ApkSet(self.apk) as z: self.check_apk(z) data = z.read(APK_AA_DIR + rel) - key = self._setting("bundle_key", b.name) if data[:7] != b"UnityFS" else None + key = self._setting("bundle_key", b.name) if not data.startswith(UNITYFS) else None _write_atomic(dst, _unityfs(data, b.name, key)) return dst @@ -354,11 +370,11 @@ def apk_bundle(self, name_contains: str) -> Path: raise KeyError(f"{name_contains!r}: {len(names)} APK bundles match") name = names[0].rsplit("/", 1)[1] dst = self.local_cache_dir() / "bundles" / name - if not (dst.exists() and dst.stat().st_size > 0): + if _cached_bundle(dst) is None: dst.parent.mkdir(parents=True, exist_ok=True) self.check_apk(z) data = z.read(names[0]) - key = self._setting("bundle_key", name) if data[:7] != b"UnityFS" else None + key = self._setting("bundle_key", name) if not data.startswith(UNITYFS) else None _write_atomic(dst, _unityfs(data, name, key)) return dst diff --git a/tests/test_catalog_corruption.py b/tests/test_catalog_corruption.py new file mode 100644 index 0000000..d47c183 --- /dev/null +++ b/tests/test_catalog_corruption.py @@ -0,0 +1,119 @@ +"""Corrupt catalog and bundle inputs must fail before becoming reusable cache entries.""" +import struct + +import pytest + +import synth +from nnnotes.addressables import (BundleKey, DYNAMIC, NONE, OFFSET_MASK, parse, + parse_header, parse_keys, parse_locations) +from nnnotes.catalog import Catalog + + +KEY = BundleKey(synth.BUNDLE_KEY, synth.BUNDLE_SEED) + + +def binary(path_ids=False): + return synth.CatalogWriter().build([ + ("Char/A", "Assets/Game/A.prefab", [1]), + ("a_01.bundle", synth.remote("a_01.bundle"), []), + ], path_ids=path_ids) + + +@pytest.mark.parametrize("bad_download", ["wrong-key", b"error response", b"", b"UnityFS"]) +def test_invalid_bundle_is_not_cached_and_can_be_retried(tmp_path, monkeypatch, bad_download): + cat = Catalog(binary(), tmp_path / "cache", cdn="https://cdn.invalid", bundle_key=KEY) + bundle, = cat.resolve("Char/A") + plain = synth.fake_bundle(bundle.name) + if bad_download == "wrong-key": + cat._settings["bundle_key"] = BundleKey(bytes(16), b"") + data = synth.encrypt_bundle(plain, bundle.name) + else: + data = bad_download + monkeypatch.setattr(cat, "_download", lambda url: data) + with pytest.raises(ValueError, match="UnityFS"): + cat.fetch(bundle) + assert cat.cached(bundle) is None + assert not list(cat.cache_dir.rglob("*.part")) + cat._settings["bundle_key"] = KEY + monkeypatch.setattr(cat, "_download", lambda url: synth.encrypt_bundle(plain, bundle.name)) + assert cat.fetch(bundle).read_bytes() == plain + + +@pytest.mark.parametrize("damaged", [b"bad cached bundle", b"UnityFS", b""]) +def test_damaged_cached_bundle_is_a_miss_and_is_replaced(tmp_path, monkeypatch, damaged): + cat = Catalog(binary(), tmp_path / "cache", cdn="https://cdn.invalid", bundle_key=KEY) + bundle, = cat.resolve("Char/A") + plain = synth.fake_bundle(bundle.name) + monkeypatch.setattr(cat, "_download", lambda url: synth.encrypt_bundle(plain, bundle.name)) + path = cat.fetch(bundle) + path.write_bytes(damaged) + assert cat.cached(bundle) is None + assert cat.fetch(bundle).read_bytes() == plain + + +def test_apk_bundle_repairs_damaged_cache(tmp_path): + import zipfile + + apk = tmp_path / "base.apk" + local = "local.bundle" + with zipfile.ZipFile(apk, "w") as archive: + archive.writestr("assets/aa/catalog.bin", synth.CatalogWriter().build([ + ("Local", synth.local(local), []), + ])) + archive.writestr("assets/aa/Android/" + local, synth.encrypt_bundle(synth.fake_bundle(local), local)) + cat = Catalog(binary(), tmp_path / "cache", apk=apk, bundle_key=KEY) + bundle, = cat.resolve("Local") + path = cat.fetch(bundle) + path.write_bytes(b"bad cached bundle") + assert cat.cached(bundle) is None + assert cat.apk_bundle("local").read_bytes() == synth.fake_bundle(local) + path.write_bytes(b"bad cached bundle") + assert cat.fetch(bundle).read_bytes() == synth.fake_bundle(local) + + +@pytest.mark.parametrize("decode", [parse, parse_header, parse_keys, parse_locations]) +@pytest.mark.parametrize("data", [b"", b"short"]) +def test_short_catalog_has_a_format_error(decode, data): + with pytest.raises(ValueError): + decode(data) + + +@pytest.mark.parametrize("decode", [parse, parse_locations]) +def test_string_cannot_extend_past_catalog(decode): + data = bytearray(binary()) + location = parse(data)[0]["offset"] + string_offset = len(data) + 4 + data += struct.pack("