#!/usr/bin/env python3 from __future__ import annotations import argparse import json import re from dataclasses import asdict, dataclass from pathlib import Path STAGE = "V8.12 generated code density" SHARED_BLOCK_MARKER = "PSPRECOMP_V812_SHARED_SCHED_BLOCK" SHARED_JUMP_MARKER = "PSPRECOMP_V812_SHARED_SCHED_JUMP" REG_MASK_MARKER = "PSPRECOMP_V812_COMPACT_REGISTRATION" RUNTIME_HELPER = "register_generated_entry_mask" EXPECTED_UNITS = 234 EXPECTED_REGISTERED_ENTRIES = 181789 EXPECTED_LOCAL_LINKS = 11299 MASK_RE = re.compile( r"static constexpr std::uint64_t (?PkEntryMasks_recomp_unit_(?P\d{4}))" r"\[(?P\d+)\] = \{\n(?P.*?)\n\};", re.DOTALL, ) REGISTER_FN_RE = re.compile( r'^[ \t]*runtime\.register_function\(0x[0-9A-Fa-f]+u,\s*&recomp_unit_\d{4},\s*"recomp_unit_\d{4}"\);\s*$', re.MULTILINE, ) REGISTER_BODY_RE = re.compile( r"void register_generated_unit_(?P\d+)\(Runtime &runtime\) \{\n" r"(?P.*?)\n\}", re.DOTALL, ) REGISTER_UNIT_CALL_RE = re.compile( r'^[ \t]*runtime\.register_generated_unit\(' r'(?P\d+)u,\s*(?P0x[0-9A-Fa-f]+)u,\s*(?P\d+)u,\s*' r'&(?Precomp_unit_(?P\d{4})),\s*&(?Precomp_unit_\d{4}_entry)\);\s*$', re.MULTILINE, ) DISPATCH_ORIGIN_RE = re.compile( r"const std::uint32_t entry_delta = local_pc - (?P0x[0-9A-Fa-f]+)u;" ) COMPACT_REG_CALL_RE = re.compile( r'runtime\.register_generated_entry_mask\((?P0x[0-9A-Fa-f]+)u,\s*' r'&(?Precomp_unit_\d{4}),\s*"(?P=fn)",\s*' r'(?PkEntryMasks_recomp_unit_\d{4}),\s*(?P\d+)u\);', re.DOTALL, ) # Exact V8.11 multi-line scheduler blocks, generated by apply_v811_perf.py. REGCACHE_LOCAL_RE = re.compile( r"(?P[ \t]*)ctx\.pc = jump_target;\n" r"(?P=i)AOT_REGCACHE_SYNC_OUT\(\);\n" r"(?P=i)// PSPRECOMP_V811_LOCAL_LINK: preserve the 256-transfer scheduler boundary, skip outer trampoline\.\n" r"(?P=i)if \(local_redispatch_rounds < 7u && rt\.continue_generated_local_dispatch\(ctx\)\) \{\n" r"(?P=i) \+\+local_redispatch_rounds;\n" r"(?P=i) AOT_REGCACHE_SYNC_IN\(\);\n" r"(?P=i) local_transfers = 0u;\n" r"(?P=i) local_pc = ctx\.pc;\n" r"(?P=i) entry_id = 0u;\n" r"(?P=i) goto LOCAL_DISPATCH;\n" r"(?P=i)\}\n" r"(?P=i)return;" ) PLAIN_LOCAL_RE = re.compile( r"(?P[ \t]*)ctx\.pc = jump_target;\n" r"(?P=i)// PSPRECOMP_V811_LOCAL_LINK: preserve the 256-transfer scheduler boundary, skip outer trampoline\.\n" r"(?P=i)if \(local_redispatch_rounds < 7u && rt\.continue_generated_local_dispatch\(ctx\)\) \{\n" r"(?P=i) \+\+local_redispatch_rounds;\n" r"(?P=i) local_transfers = 0u;\n" r"(?P=i) local_pc = ctx\.pc;\n" r"(?P=i) entry_id = 0u;\n" r"(?P=i) goto LOCAL_DISPATCH;\n" r"(?P=i)\}\n" r"(?P=i)return;" ) @dataclass class UnitStats: path: str bytes_before: int bytes_after: int entry_count: int registration_lines_removed: int local_blocks_shared: int had_regcache: bool def read(path: Path) -> str: return path.read_text(encoding="utf-8", errors="strict") def write_if_changed(path: Path, old: str, new: str) -> bool: if old == new: return False tmp = path.with_suffix(path.suffix + ".v812tmp") tmp.write_text(new, encoding="utf-8", newline="\n") tmp.replace(path) return True def require_once(text: str, needle: str, path: Path, description: str) -> None: count = text.count(needle) if count != 1: raise RuntimeError(f"{path}: expected one {description}, found {count}") def parse_masks(text: str, path: Path) -> tuple[str, int, list[int], int]: m = MASK_RE.search(text) if not m: raise RuntimeError(f"{path}: V8.11 compact entry masks not found") declared = int(m.group("count")) values = [int(x, 16) for x in re.findall(r"0x([0-9A-Fa-f]{16})ull", m.group("body"))] if len(values) != declared: raise RuntimeError(f"{path}: compact mask declared {declared}, parsed {len(values)}") entries = sum(x.bit_count() for x in values) return m.group("name"), declared, values, entries def patch_runtime_hpp(path: Path) -> bool: old = read(path) if RUNTIME_HELPER in old: return False anchor = ( " void register_generated_unit(std::uint32_t unit_index, std::uint32_t unit_address,\n" " std::uint32_t unit_span, RecompiledFunction function,\n" " RecompiledEntryFunction entry_function = nullptr);\n" ) require_once(old, anchor, path, "register_generated_unit declaration") addition = anchor + ( " // V8.12: register all valid per-PC entries from the compact 64-slot\n" " // occupancy masks already emitted for generated-unit dispatch. This\n" " // preserves the exact legacy function/direct-chain tables while avoiding\n" " // hundreds of source-level registration call sites per translation unit.\n" " void register_generated_entry_mask(std::uint32_t unit_address,\n" " RecompiledFunction function,\n" " std::string_view name,\n" " const std::uint64_t *entry_masks,\n" " std::size_t group_count);\n" ) return write_if_changed(path, old, old.replace(anchor, addition, 1)) def patch_runtime_cpp(path: Path) -> bool: old = read(path) text = old if "#include \n" not in text: anchor = "#include \n" require_once(text, anchor, path, " include") text = text.replace(anchor, anchor + "#include \n", 1) if f"Runtime::{RUNTIME_HELPER}" not in text: # Insert immediately after register_generated_unit(), before invoke_isolated_aot(). anchor = "\nbool Runtime::invoke_isolated_aot(std::uint32_t address, AllegrexContext &ctx) {\n" require_once(text, anchor, path, "invoke_isolated_aot implementation") impl = ( "\n// PSPRECOMP_V812_COMPACT_REGISTRATION\n" "void Runtime::register_generated_entry_mask(std::uint32_t unit_address,\n" " RecompiledFunction function,\n" " std::string_view name,\n" " const std::uint64_t *entry_masks,\n" " std::size_t group_count) {\n" " if (function == nullptr || entry_masks == nullptr || group_count == 0u) return;\n" " for (std::size_t group = 0; group < group_count; ++group) {\n" " std::uint64_t bits = entry_masks[group];\n" " while (bits != 0u) {\n" " const auto bit = static_cast(std::countr_zero(bits));\n" " const auto slot = static_cast(group * 64u) + bit;\n" " register_function(unit_address + slot * 4u, function, std::string(name));\n" " bits &= bits - 1u;\n" " }\n" " }\n" "}\n" ) text = text.replace(anchor, impl + anchor, 1) return write_if_changed(path, old, text) def insert_shared_scheduler_block(text: str, path: Path, has_regcache: bool) -> str: if SHARED_BLOCK_MARKER in text: return text # First guest label follows the LOCAL_DISPATCH switch and is at outer function scope. dispatch = text.find("LOCAL_DISPATCH:\n") if dispatch < 0: raise RuntimeError(f"{path}: LOCAL_DISPATCH not found") m = re.search(r"(?m)^L_[0-9A-Fa-f]+:\n", text[dispatch:]) if not m: raise RuntimeError(f"{path}: first guest label after LOCAL_DISPATCH not found") pos = dispatch + m.start() lines = [ f"// {SHARED_BLOCK_MARKER}: one scheduler-exact slow boundary per unit.", "LOCAL_SCHED_BOUNDARY:", " ctx.pc = jump_target;", ] if has_regcache: lines.append(" AOT_REGCACHE_SYNC_OUT();") lines.extend([ " if (local_redispatch_rounds < 7u && rt.continue_generated_local_dispatch(ctx)) {", " ++local_redispatch_rounds;", ]) if has_regcache: lines.append(" AOT_REGCACHE_SYNC_IN();") lines.extend([ " local_transfers = 0u;", " local_pc = ctx.pc;", " entry_id = 0u;", " goto LOCAL_DISPATCH;", " }", " return;", "", ]) return text[:pos] + "\n".join(lines) + text[pos:] def share_local_scheduler_blocks(text: str, path: Path) -> tuple[str, int]: if SHARED_BLOCK_MARKER in text: return text, text.count(SHARED_JUMP_MARKER) has_regcache = "#define AOT_REGCACHE_SYNC_OUT()" in text count = 0 def repl(m: re.Match[str]) -> str: nonlocal count count += 1 i = m.group("i") return f"{i}// {SHARED_JUMP_MARKER}\n{i}goto LOCAL_SCHED_BOUNDARY;" if has_regcache: text = REGCACHE_LOCAL_RE.sub(repl, text) # A register-cached unit must not contain an unwrapped V8.11 block. if "PSPRECOMP_V811_LOCAL_LINK" in text: raise RuntimeError(f"{path}: unrecognized V8.11 local-link form remains in regcache unit") else: text = PLAIN_LOCAL_RE.sub(repl, text) if "PSPRECOMP_V811_LOCAL_LINK" in text: raise RuntimeError(f"{path}: unrecognized V8.11 local-link form remains in plain unit") if count: text = insert_shared_scheduler_block(text, path, has_regcache) return text, count def compact_registration(text: str, path: Path, mask_name: str, group_count: int, entry_count: int) -> tuple[str, int]: # Already compacted. if REG_MASK_MARKER in text: return text, entry_count matches = list(REGISTER_BODY_RE.finditer(text)) if len(matches) != 1: raise RuntimeError(f"{path}: expected one register_generated_unit_N body, found {len(matches)}") m = matches[0] body = m.group("body") unit_call = REGISTER_UNIT_CALL_RE.search(body) if unit_call is None: raise RuntimeError(f"{path}: register_generated_unit call not found") unit_suffix = unit_call.group("unit") fn = unit_call.group("fn") bucket = unit_call.group("bucket") unit_base = unit_call.group("base") origin_match = DISPATCH_ORIGIN_RE.search(text) if origin_match is None: raise RuntimeError(f"{path}: compact-dispatch entry origin not found") entry_origin = origin_match.group("origin") if mask_name != f"kEntryMasks_{fn}": raise RuntimeError(f"{path}: mask/function mismatch {mask_name} vs {fn}") registration_lines = REGISTER_FN_RE.findall(body) reg_count = len(registration_lines) if reg_count != entry_count: raise RuntimeError( f"{path}: registration lines={reg_count} but compact masks contain {entry_count} entries" ) # Require the old body to consist only of the unit registration plus exact per-PC registrations. scrubbed = REGISTER_FN_RE.sub("", body) scrubbed = REGISTER_UNIT_CALL_RE.sub("", scrubbed) if scrubbed.strip(): raise RuntimeError(f"{path}: unexpected code in generated registration body: {scrubbed.strip()[:120]!r}") unit_call_line = unit_call.group(0).strip() new_body = ( f"void register_generated_unit_{m.group('bucket')}(Runtime &runtime) {{\n" f" {unit_call_line}\n" f" // {REG_MASK_MARKER}: same entry set, looped once in Runtime instead of source-expanded calls.\n" # IMPORTANT: kEntryMasks is relative to the compact dispatch origin, # which is not always the 16 KiB generated-unit base (some units begin # at base+4). Using unit_base shifts every registered PC and breaks boot. f" runtime.register_generated_entry_mask({entry_origin}u, &{fn}, \"{fn}\",\n" f" {mask_name}, {group_count}u);\n" f"}}" ) text = text[:m.start()] + new_body + text[m.end():] return text, reg_count def transform_unit(path: Path) -> UnitStats: old = read(path) text = old mask_name, group_count, _masks, entry_count = parse_masks(text, path) had_regcache = "#define AOT_REGCACHE_SYNC_OUT()" in text text, local_count = share_local_scheduler_blocks(text, path) text, registration_count = compact_registration(text, path, mask_name, group_count, entry_count) write_if_changed(path, old, text) return UnitStats( path=path.name, bytes_before=len(old.encode("utf-8")), bytes_after=len(text.encode("utf-8")), entry_count=entry_count, registration_lines_removed=registration_count if REG_MASK_MARKER not in old else 0, local_blocks_shared=local_count if SHARED_BLOCK_MARKER not in old else 0, had_regcache=had_regcache, ) def validate(root: Path) -> dict: profile = root / "profiles" / "vcs" generated = profile / "generated" units = sorted(generated.glob("generated_unit_*.cpp")) if len(units) != EXPECTED_UNITS: raise RuntimeError(f"expected {EXPECTED_UNITS} generated units, found {len(units)}") runtime_hpp = read(root / "include" / "psprecomp" / "runtime.hpp") runtime_cpp = read(root / "src" / "runtime.cpp") if RUNTIME_HELPER not in runtime_hpp or f"Runtime::{RUNTIME_HELPER}" not in runtime_cpp: raise RuntimeError("V8.12 compact registration Runtime helper is missing") total_entries = 0 total_jumps = 0 shared_units = 0 compact_registration_units = 0 total_bytes = 0 for path in units: text = read(path) _mask_name, _groups, _masks, entries = parse_masks(text, path) total_entries += entries total_bytes += len(text.encode("utf-8")) jumps = text.count(SHARED_JUMP_MARKER) total_jumps += jumps if jumps: if text.count(SHARED_BLOCK_MARKER) != 1: raise RuntimeError(f"{path}: shared jumps present without exactly one shared block") if text.count("LOCAL_SCHED_BOUNDARY:") != 1: raise RuntimeError(f"{path}: expected exactly one LOCAL_SCHED_BOUNDARY label") shared_units += 1 elif SHARED_BLOCK_MARKER in text: raise RuntimeError(f"{path}: shared block exists with no shared jumps") if "PSPRECOMP_V811_LOCAL_LINK" in text: raise RuntimeError(f"{path}: old duplicated V8.11 local-link block remains") if text.count(REG_MASK_MARKER) != 1: raise RuntimeError(f"{path}: compact registration marker missing/duplicated") if text.count("runtime.register_generated_entry_mask(") != 1: raise RuntimeError(f"{path}: expected one compact registration call") origin_match = DISPATCH_ORIGIN_RE.search(text) compact_match = COMPACT_REG_CALL_RE.search(text) if origin_match is None or compact_match is None: raise RuntimeError(f"{path}: compact registration/origin parse failed") dispatch_origin = int(origin_match.group("origin"), 16) registration_origin = int(compact_match.group("origin"), 16) if registration_origin != dispatch_origin: raise RuntimeError( f"{path}: compact registration origin {registration_origin:#010x} " f"!= dispatch origin {dispatch_origin:#010x}" ) if compact_match.group("mask") != _mask_name or int(compact_match.group("groups")) != _groups: raise RuntimeError(f"{path}: compact registration mask/group mismatch") # Prove that every mask bit resolves to exactly the same PC represented # by the unit's switch table. This catches any future origin drift. switch_pcs = { int(x, 16) for x in re.findall(r"case\s+\d+u:\s+goto\s+L_([0-9A-Fa-f]{8});", text) } mask_pcs = { dispatch_origin + (group * 64 + bit) * 4 for group, mask in enumerate(_masks) for bit in range(64) if (mask >> bit) & 1 } if switch_pcs != mask_pcs: missing = sorted(switch_pcs - mask_pcs)[:8] extra = sorted(mask_pcs - switch_pcs)[:8] raise RuntimeError( f"{path}: compact mask/switch PC mismatch " f"missing={[hex(x) for x in missing]} extra={[hex(x) for x in extra]}" ) # Per-PC generated registration calls must be gone from this unit's registration function. rm = REGISTER_BODY_RE.search(text) if rm is None: raise RuntimeError(f"{path}: generated registration body missing after V8.12") if REGISTER_FN_RE.search(rm.group("body")): raise RuntimeError(f"{path}: old per-PC registration calls remain") compact_registration_units += 1 if total_entries != EXPECTED_REGISTERED_ENTRIES: raise RuntimeError( f"entry-mask population changed: expected {EXPECTED_REGISTERED_ENTRIES}, found {total_entries}" ) if total_jumps != EXPECTED_LOCAL_LINKS: raise RuntimeError( f"local-link population changed: expected {EXPECTED_LOCAL_LINKS}, found {total_jumps}" ) return { "units": len(units), "entry_count": total_entries, "shared_local_jumps": total_jumps, "shared_local_units": shared_units, "compact_registration_units": compact_registration_units, "generated_cpp_bytes": total_bytes, } def main() -> int: parser = argparse.ArgumentParser(description=STAGE) parser.add_argument("root", nargs="?", default=".", help="PSPRecomp repository root") parser.add_argument("--check", action="store_true", help="validate only") args = parser.parse_args() root = Path(args.root).resolve() profile = root / "profiles" / "vcs" generated = profile / "generated" if not (root / "CMakeLists.txt").is_file() or not generated.is_dir(): raise RuntimeError(f"not a PSPRecomp VCS tree: {root}") manifest_path = generated / "v812_code_density_manifest.json" if args.check: summary = validate(root) print( "V8.12 CHECK OK: " f"units={summary['units']} entries={summary['entry_count']} " f"shared_jumps={summary['shared_local_jumps']} " f"generated_cpp={summary['generated_cpp_bytes'] / (1024*1024):.2f} MiB" ) return 0 units = sorted(generated.glob("generated_unit_*.cpp")) if len(units) != EXPECTED_UNITS: raise RuntimeError(f"expected {EXPECTED_UNITS} generated units, found {len(units)}") before_total = sum(p.stat().st_size for p in units) previous = {} if manifest_path.is_file(): try: previous = json.loads(manifest_path.read_text(encoding="utf-8")) except Exception: previous = {} patch_runtime_hpp(root / "include" / "psprecomp" / "runtime.hpp") patch_runtime_cpp(root / "src" / "runtime.cpp") stats = [transform_unit(path) for path in units] after_total = sum(p.stat().st_size for p in units) baseline = int(previous.get("baseline_generated_cpp_bytes", before_total)) if baseline < after_total and previous.get("stage") != STAGE: baseline = before_total summary = validate(root) manifest = { "stage": STAGE, "baseline_generated_cpp_bytes": baseline, "generated_cpp_bytes_before_this_run": before_total, "generated_cpp_bytes_after": after_total, "source_bytes_saved_this_run": before_total - after_total, "source_bytes_saved_from_baseline": baseline - after_total, "registered_entries_preserved": summary["entry_count"], "per_pc_registration_calls_collapsed_this_run": sum(s.registration_lines_removed for s in stats), "scheduler_blocks_shared_this_run": sum(s.local_blocks_shared for s in stats), "shared_local_jumps": summary["shared_local_jumps"], "shared_local_units": summary["shared_local_units"], "compact_registration_units": summary["compact_registration_units"], "units": [asdict(s) for s in stats], } manifest_path.write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8") print( "V8.12 code-density pass applied: " f"units={len(stats)} entries={summary['entry_count']} " f"registrations_collapsed={sum(s.registration_lines_removed for s in stats)} " f"shared_boundaries={sum(s.local_blocks_shared for s in stats)} " f"generated_cpp={after_total / (1024*1024):.2f} MiB " f"saved={(before_total-after_total) / (1024*1024):.2f} MiB this run " f"cumulative={(baseline-after_total) / (1024*1024):.2f} MiB" ) return 0 if __name__ == "__main__": raise SystemExit(main())