refactor: extract companion extras and the kept filter

This commit is contained in:
Abdessamad Derraz committed 2026-10-06 07:03:59 +02:00
1 parent 97ccf9027b
commit 3dce7b5f09
2 files changed
+61 -35

No files matched your search

+51 -33
View File
@@ -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
+10 -2
View File
@@ -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__":