diff --git a/scripts/packextras.py b/scripts/packextras.py index 2a7d05b2..8b108c30 100644 --- a/scripts/packextras.py +++ b/scripts/packextras.py @@ -266,30 +266,44 @@ def _agnostic_scan_extras( "agnostic_scan": True, } ) - for companion_sha1, companion in _companions( - path, scan_name, _companion_extensions(profile, fname), files_db - ): - companion_dest = f"{folder}/{companion['name']}" if folder else companion["name"] - companion_full = ( - f"{extras_prefix}/{companion_dest}" if extras_prefix else companion_dest - ) - if companion_full in seen_dests: - continue - seen_dests.add(companion_full) - extras.append({ - "name": companion["name"], - "destination": companion_dest, - "sha1": companion_sha1, - "required": False, - "hle_fallback": False, - "source_emulator": profile.get("emulator", emu_name), - "source_profile": emu_name, - "source_system": f.get("system"), - "source_systems": list(profile.get("systems", [])), - "region": f.get("region"), - "variant_group": f.get("variant_group"), - "agnostic_scan": True, - }) + extras.extend(_companion_extras( + _companions(path, scan_name, _companion_extensions(profile, fname), files_db), + folder, extras_prefix, seen_dests, profile, emu_name, f, + )) + return extras + + +def _companion_extras( + companions: list[tuple[str, dict]], + folder: str, + extras_prefix: str, + seen_dests: set[str], + profile: dict, + emu_name: str, + seed: dict, +) -> list[dict]: + """Extras for the companions of one scanned image, placed beside it.""" + extras: list[dict] = [] + for companion_sha1, companion in companions: + companion_dest = f"{folder}/{companion['name']}" if folder else companion["name"] + companion_full = f"{extras_prefix}/{companion_dest}" if extras_prefix else companion_dest + if companion_full in seen_dests: + continue + seen_dests.add(companion_full) + extras.append({ + "name": companion["name"], + "destination": companion_dest, + "sha1": companion_sha1, + "required": False, + "hle_fallback": False, + "source_emulator": profile.get("emulator", emu_name), + "source_profile": emu_name, + "source_system": seed.get("system"), + "source_systems": list(profile.get("systems", [])), + "region": seed.get("region"), + "variant_group": seed.get("variant_group"), + "agnostic_scan": True, + }) return extras @@ -615,6 +629,13 @@ def _extra_system_ids(extra: dict) -> list[str]: return [str(explicit)] return [str(value) for value in extra.get("source_systems", []) if value] +def _kept(entries: list[dict], required_only: bool) -> list[dict]: + """The entries a pack keeps: optional ones go under --required-only.""" + if not required_only: + return entries + return [e for e in entries if e.get("required") is not False] + + def platform_region_groups( config: dict, systems: dict, @@ -639,11 +660,9 @@ def platform_region_groups( groups: dict[str, list[tuple[str, str]]] = {} for sys_id, system in systems.items(): members = groups.setdefault(sys_id, []) - for file_entry in system.get("files", []): - if required_only and file_entry.get("required") is False: - # Decided over what the pack keeps: an optional winner - # removed afterwards left its slot empty. - continue + # Decided over what the pack keeps: an optional winner removed + # afterwards left its slot empty. + for file_entry in _kept(system.get("files", []), required_only): dest = sanitize_pack_path( file_entry.get("destination", file_entry.get("name", "")) ) @@ -654,7 +673,7 @@ def platform_region_groups( if not include_extras or db is None: return groups, extra_dests - for extra in _collect_emulator_extras( + extras = _collect_emulator_extras( config, emulators_dir, db, @@ -663,9 +682,8 @@ def platform_region_groups( emu_profiles, target_cores=target_cores, include_all=include_all, - ): - if required_only and extra.get("required") is False: - continue + ) + for extra in _kept(extras, required_only): dest = sanitize_pack_path(extra.get("destination", extra.get("name", ""))) if not dest: continue diff --git a/tests/test_region_required_only.py b/tests/test_region_required_only.py index f789e1e4..40ab2f6f 100644 --- a/tests/test_region_required_only.py +++ b/tests/test_region_required_only.py @@ -8,6 +8,7 @@ the optional winner: the slot shipped empty. from __future__ import annotations +import ast import sys import unittest from pathlib import Path @@ -30,8 +31,15 @@ class RequiredOnlyBeforeRegion(unittest.TestCase): self.assertEqual(groups["vic20"], [("Vic20/kernel-ntsc.bin", "kernel-ntsc.bin")]) def test_every_caller_passes_it(self): - source = (REPO_ROOT / "scripts" / "generate_pack.py").read_text(encoding="utf-8") - self.assertEqual(source.count("required_only=required_only,\n )"), 2) + tree = ast.parse((REPO_ROOT / "scripts" / "generate_pack.py").read_text(encoding="utf-8")) + calls = [ + node for node in ast.walk(tree) + if isinstance(node, ast.Call) + and getattr(node.func, "id", None) == "platform_region_groups" + ] + self.assertGreaterEqual(len(calls), 2) + for call in calls: + self.assertIn("required_only", {kw.arg for kw in call.keywords}) if __name__ == "__main__":