"""Deterministic conflict resolution for import pull requests. Backs ``.github/workflows/auto-rebase.yml``. When a pull request merges, every other open import PR conflicts in exactly two files, and in both the resolution is mechanical rather than editorial: ``LeanPool.lean`` The ``mk_all`` index is a sorted list of ``import LeanPool.X`` lines, one per project aggregate after migration, regenerated rather than merged -- and, because it is derived purely from the file tree, without needing a Lean toolchain. ``LeanPool/projects.yml`` Take the merged base's registry and re-append the cards this branch added. Cards are moved as verbatim text blocks, never re-serialised: round-tripping 141 cards through a YAML dumper would reformat every one of them and bury the real change. Anything else in conflict is a genuine content overlap and is left alone for a human. This module only computes file contents; the workflow decides what to do with them. """ from __future__ import annotations import argparse import logging import re import sys from pathlib import Path import yaml from lean_pool.exposition.source_text import code_view from lean_pool.indexes import main as regenerate_indexes from lean_pool.indexes import render_project_indexes, requires_project_roots logger = logging.getLogger(__name__) REGISTRY = "LeanPool/projects.yml" INDEX = "LeanPool.lean" # The only conflicts this module claims to resolve. RESOLVABLE = frozenset({INDEX, REGISTRY}) def _uses_module_system(source: str) -> bool: """Detect a module header in either current version of a conflicted index.""" conflict = re.compile( r"^<<<<<<<[^\n]*\n(?P.*?)" r"(?:^\|{7}[^\n]*\n.*?)?^=======\n" r"(?P.*?)^>>>>>>>[^\n]*(?:\n|$)", re.MULTILINE | re.DOTALL, ) # Keep shared text in both versions so comments spanning a conflict remain # comments, while a malformed comment on one side cannot mask the other. return any( re.search( r"^\s*module(?:\s|$)", code_view(conflict.sub(lambda match: match[side], source)), re.MULTILINE, ) for side in ("ours", "theirs") ) def render_index(root: Path) -> str: """Regenerate the ``mk_all`` index from the Lean files on disk.""" if requires_project_roots(root): return render_project_indexes(root)[root / INDEX] pool = root / "LeanPool" modules = sorted( "LeanPool." + str(path.relative_to(pool)).removesuffix(".lean").replace("/", ".") for path in pool.rglob("*.lean") ) index = root / INDEX existing = index.read_text(encoding="utf-8") if index.exists() else "" uses_modules = _uses_module_system(existing) header = "module -- shake: keep-all --deprecated_module: ignore\n\n" prefix = "public " if uses_modules else "" return (header if uses_modules else "") + "".join( f"{prefix}import {module}\n" for module in modules ) def _project_nodes(text: str) -> yaml.SequenceNode: """Read card boundaries from YAML rather than assuming a first key.""" root = yaml.compose(text, Loader=yaml.SafeLoader) if not isinstance(root, yaml.MappingNode) or len(root.value) != 1: raise ValueError("Expected a registry containing only projects") key, projects = root.value[0] if key.value != "projects" or not isinstance(projects, yaml.SequenceNode): raise ValueError("Expected projects to be a sequence") if projects.flow_style and projects.value: raise ValueError("Project cards must use a block sequence") if any(not isinstance(card, yaml.MappingNode) for card in projects.value): raise ValueError("Each project card must be a mapping") return projects def _card_slug(card: yaml.MappingNode) -> str: """Require one unambiguous string slug, in any mapping position.""" slugs = [value for key, value in card.value if key.value == "slug"] if ( len(slugs) != 1 or not isinstance(slugs[0], yaml.ScalarNode) or slugs[0].tag != "tag:yaml.org,2002:str" or not slugs[0].value.strip() ): raise ValueError("Each project card must have exactly one nonempty string slug") return slugs[0].value def split_cards(text: str) -> tuple[str, list[tuple[str, str]]]: """Split a registry into its header and its cards, as verbatim text. Returns ``(header, [(slug, block)])`` where concatenating the header and every block reproduces ``text`` exactly. """ projects = _project_nodes(text) nodes = projects.value entries = [ token for token in yaml.scan(text, Loader=yaml.SafeLoader) if isinstance(token, yaml.tokens.BlockEntryToken) and token.start_mark.column == projects.start_mark.column ] starts = [text.rfind("\n", 0, token.start_mark.index) + 1 for token in entries] if len(starts) != len(nodes): raise ValueError("Expected one sequence item per project card") if not starts: return text, [] header = text[: starts[0]] cards: list[tuple[str, str]] = [] for index, (start, node) in enumerate(zip(starts, nodes, strict=True)): end = starts[index + 1] if index + 1 < len(starts) else len(text) slug = _card_slug(node) if any(existing == slug for existing, _ in cards): raise ValueError(f"Duplicate project slug: {slug}") cards.append((slug, text[start:end])) return header, cards def merge_registry(base: str, ours: str, theirs: str) -> str: """Three-way merge the registry by card. ``ours`` is the updated base branch, ``theirs`` the pull request. The result is ``ours`` plus every card the pull request added, appended in the order the pull request had them. Cards are never reordered or reformatted, so the diff shows only the additions. """ base_slugs = {slug for slug, _ in split_cards(base)[1]} header, our_cards = split_cards(ours) our_slugs = {slug for slug, _ in our_cards} added = [ (slug, block) for slug, block in split_cards(theirs)[1] if slug not in base_slugs and slug not in our_slugs ] if not added: return ours merged = header + "".join(block for _, block in our_cards) # A registry whose last card lacks a trailing newline would otherwise # run into the first appended card. if merged and not merged.endswith("\n"): merged += "\n" result = merged + "".join(block for _, block in added) # Reject incompatible layouts instead of pushing a malformed registry. split_cards(result) return result def resolvable(conflicts: list[str]) -> bool: """Whether every conflicted path is one this module can resolve.""" return bool(conflicts) and all( path in RESOLVABLE or re.fullmatch(r"LeanPool/[^/]+/Imports\.lean", path) for path in conflicts ) def _command_index(args: argparse.Namespace) -> int: """Rewrite the index from the working tree.""" root = args.repo.resolve() if requires_project_roots(root): return regenerate_indexes(["--repo", str(root), "--lib", "LeanPool"]) (root / INDEX).write_text(render_index(root), encoding="utf-8") logger.info("regenerated %s", INDEX) return 0 def _command_registry(args: argparse.Namespace) -> int: """Three-way merge the registry from three revisions on disk.""" merged = merge_registry( args.base.read_text(encoding="utf-8"), args.ours.read_text(encoding="utf-8"), args.theirs.read_text(encoding="utf-8"), ) (args.repo.resolve() / REGISTRY).write_text(merged, encoding="utf-8") logger.info("merged %s", REGISTRY) return 0 def _command_resolvable(args: argparse.Namespace) -> int: """Exit 0 when every conflicted path is mechanically resolvable.""" conflicts = [line.strip() for line in args.conflicts.read_text().splitlines()] conflicts = [path for path in conflicts if path] if resolvable(conflicts): logger.info("all conflicts are mechanically resolvable") return 0 unresolvable = sorted(set(conflicts) - RESOLVABLE) logger.error("conflicts need a human: %s", ", ".join(unresolvable) or "none") return 1 def _parse_args(argv: list[str] | None) -> argparse.Namespace: """Parse command-line arguments.""" parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) subparsers = parser.add_subparsers(dest="command", required=True) index = subparsers.add_parser("index", help="regenerate LeanPool.lean") index.add_argument("--repo", type=Path, default=Path(".")) index.set_defaults(func=_command_index) registry = subparsers.add_parser("registry", help="three-way merge projects.yml") registry.add_argument("--repo", type=Path, default=Path(".")) registry.add_argument("--base", type=Path, required=True) registry.add_argument("--ours", type=Path, required=True) registry.add_argument("--theirs", type=Path, required=True) registry.set_defaults(func=_command_registry) check = subparsers.add_parser("resolvable", help="are these conflicts mechanical?") check.add_argument("--conflicts", type=Path, required=True) check.set_defaults(func=_command_resolvable) return parser.parse_args(argv) def main(argv: list[str] | None = None) -> int: """Dispatch a subcommand; return a process exit code.""" logging.basicConfig(level=logging.INFO, format="%(message)s") args = _parse_args(argv) return int(args.func(args)) if __name__ == "__main__": sys.exit(main())