refactor: move flag refusals out of the main functions

This commit is contained in:
Abdessamad Derraz committed 2026-10-06 03:08:35 +02:00
1 parent a2e5d73bf5
commit ab75123320
4 files changed
+136 -98

No files matched your search

+58 -34
View File
@@ -11,6 +11,7 @@ from __future__ import annotations
import argparse
import contextlib
import difflib
import itertools
import json
import os
import posixpath
@@ -1256,6 +1257,28 @@ def detect_pinned_tag(
return None
def _resolve_project_words(refs: list, tree_order) -> list:
"""Point `project file` parts at the declared tree that vendors them.
tree_order() yields the trees to search, pin before HEAD; it is only
called for a part that needs it, since reading a HEAD tree costs a
request. Parts no tree carries stay external.
"""
resolved: dict[str, str | None] = {}
out = []
for display, ref, tokens, hashes, parts, citation in refs:
rewritten = []
for part in parts:
if is_external_citation(part.path):
if part.path not in resolved:
resolved[part.path] = resolve_project_word(part.path, tree_order())
if resolved[part.path]:
part = RefPart(resolved[part.path], part.start, part.end, part.raw)
rewritten.append(part)
out.append((display, ref, tokens, hashes, rewritten, citation))
return out
def build_report(
name: str, profile: dict, cache_dir: str, offline: bool = False
) -> ProfileReport:
@@ -1545,26 +1568,13 @@ def build_report(
candidates = near
return None, candidates
def _tree_order():
for view in views:
yield _pin_tree_for(view)
for view in views:
yield _context_for(view)[1]
resolved_words: dict[str, str | None] = {}
for index, (display, ref, tokens, hashes, parts, citation) in enumerate(refs):
rewritten = []
for part in parts:
if is_external_citation(part.path):
if part.path not in resolved_words:
resolved_words[part.path] = resolve_project_word(
part.path, _tree_order()
)
found = resolved_words[part.path]
if found:
part = RefPart(found, part.start, part.end, part.raw)
rewritten.append(part)
refs[index] = (display, ref, tokens, hashes, rewritten, citation)
# Pin trees before HEAD trees, each read only when reached.
refs = _resolve_project_words(
refs,
lambda: itertools.chain(
map(_pin_tree_for, views), map(lambda v: _context_for(v)[1], views)
),
)
lines_cache: dict[tuple[str, str, int | None], list[str] | None] = {}
@@ -2571,6 +2581,25 @@ def _realign_part(
return "skip", f"{part.path} absent at the writing revision"
def _writing_pairs(
document: dict, repos: list, revisions: list[tuple[str, dict]], intro_sha: str
) -> tuple[list, list[str]]:
"""(repo, writing pin, current pin) per declared repository, and the
repositories whose writing pin the history cannot single out."""
pairs = []
ambiguous: list[str] = []
for pin_field, repo in repos:
current = document.get(f"{pin_field}_commit")
if not isinstance(current, str) or not current:
continue
written = _writing_pins(revisions, intro_sha, f"{pin_field}_commit")
if len(written) > 1:
ambiguous.append(f"{pin_field}_commit {' -> '.join(written)}")
elif written and written[0] != current:
pairs.append((repo, written[0], current))
return pairs, ambiguous
def realign_prose(
path: Path, cache_dir: str, offline: bool = False, dry_run: bool = False
) -> list[str]:
@@ -2629,17 +2658,7 @@ def realign_prose(
if intro_sha is None:
# The scalar is not committed yet: written now, under this pin.
continue
pairs = []
ambiguous: list[str] = []
for pin_field, repo in repos:
current = document.get(f"{pin_field}_commit")
if not isinstance(current, str) or not current:
continue
written = _writing_pins(revisions, intro_sha, f"{pin_field}_commit")
if len(written) > 1:
ambiguous.append(f"{pin_field}_commit {' -> '.join(written)}")
elif written and written[0] != current:
pairs.append((repo, written[0], current))
pairs, ambiguous = _writing_pairs(document, repos, revisions, intro_sha)
if ambiguous:
# Moving from the wrong one rewrites a correct citation onto
# someone else's code, with nothing to show it happened.
@@ -2985,9 +3004,8 @@ def _print_triage(args, selected: dict, reports: list[ProfileReport]) -> None:
print(f"{score:5d} {report.name:30s} {state}")
def main() -> None:
"""Build one report per selected profile, then apply writes and output."""
args = build_parser().parse_args()
def _refuse_flag_combinations(args) -> None:
"""Exit on a flag the chosen mode would not apply, before any request."""
if args.accept_changed and not args.emulator:
print(
"--accept-changed applies to one profile at a time: its diff has "
@@ -3052,6 +3070,12 @@ def main() -> None:
file=sys.stderr,
)
raise SystemExit(1)
def main() -> None:
"""Build one report per selected profile, then apply writes and output."""
args = build_parser().parse_args()
_refuse_flag_combinations(args)
profiles = load_emulator_profiles(args.emulators_dir, skip_aliases=False)
selected = select_profiles(profiles, args)
_check_quota(len(selected), args.offline)