Coverage for src/renaissance/syntax_tree/match_finder.py: 98%
222 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-10-02 13:12 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-10-02 13:12 +0000
1"""AI: Pattern matching engine that finds AST nodes matching a given pattern."""
3from collections.abc import Iterable, Sequence
4from typing import Self
6from renaissance.utils.ast_utils import use_dollar
8from .node_protocol import NodeProtocol
9from .pattern_kind import PatternKind
10from .semantic_kind import SemanticKind
12IRRELEVANT_PROPS = {"macro_expansion", "start_point", "end_point", "source_code", "location", "type"}
14MIS_MATCH = -12
15INCOMPLETE_MATCH = -11
16_TOP_LEVEL_KINDS = {"Module", "TRANSLATION_UNIT"}
19def pattern_kind(node: NodeProtocol) -> PatternKind | None:
20 """AI: Return the pattern kind (MATCH_ONE/MATCH_ALL) of a pattern placeholder node, or None if not one."""
21 value = getattr(node, "pattern_kind", None)
22 if value is not None:
23 return value
24 if node.parser_kind in {"MatchOne", "_MatchOne__"}:
25 return PatternKind.MATCH_ONE
26 if node.parser_kind in {"MatchAll", "_MatchAll__"}:
27 return PatternKind.MATCH_ALL
28 return None
31def node_kinds_match(source: NodeProtocol, pattern: NodeProtocol) -> bool:
32 """AI: Return True if source and pattern nodes have the same semantic kind (or parser kind as fallback)."""
33 if (
34 source.semantic_kind is not None
35 and pattern.semantic_kind is not None
36 and source.semantic_kind is not SemanticKind.NODE
37 and pattern.semantic_kind is not SemanticKind.NODE
38 ):
39 return source.semantic_kind == pattern.semantic_kind
40 return source.parser_kind == pattern.parser_kind
43class Variant:
44 """AI: Track one candidate pattern-match state (bound expansions, greedy position) during matching."""
46 def __init__(
47 self,
48 index: int,
49 exp: dict[str, Sequence[NodeProtocol]],
50 greedy: str | None,
51 expansion_start: int,
52 end_index: int = INCOMPLETE_MATCH,
53 ):
54 """AI: Track one candidate pattern-match state (bound expansions, greedy position) during matching."""
55 self.exp: dict[str, Sequence[NodeProtocol]] = exp
56 self.index: int = index
57 self.greedy: str | None = greedy
58 self.end_index: int = end_index
59 self.expansion_start: int = expansion_start
61 def reset_greedy(self):
62 """AI: Clear the current greedy-expansion tracking state."""
63 self.greedy = None
64 self.expansion_start = -1
66 def close_greedy(self, key, nodes, start, end):
67 """Store a completed greedy expansion and reset greedy state."""
68 value = nodes[start:end]
69 self.exp[key] = value
70 self.reset_greedy()
72 def fork(self) -> Variant:
73 """Return a copy of this variant at the same position (for backtracking)."""
74 return Variant(self.index, self.exp.copy(), self.greedy, self.expansion_start)
77class PatternMatch:
78 """AI: Represent a successful match of a pattern against a sequence of AST nodes."""
80 def __init__(self, nodes, expansions, patterns):
81 """AI: Represent a successful match of a pattern against a sequence of AST nodes."""
82 self.nodes = nodes
83 self.expansions = expansions
84 self.patterns = patterns
86 def __str__(self):
87 """AI: Return the newline-joined signatures of the matched nodes."""
88 return "\n".join(node.signature for node in self.nodes)
90 @property
91 def signature(self):
92 """AI: Return the newline-joined signatures of the matched nodes."""
93 return str(self)
95 def __getitem__(self, key):
96 """AI: Return the newline-joined expansion text for the given placeholder key."""
97 return "\n".join(node.signature if isinstance(node, NodeProtocol) else node for node in self.expansions[key])
99 def match_referenced_by(self, patterns: Sequence[list], recursive: bool = True) -> Sequence[Self]:
100 """AI: Return matches of patterns against the nodes that reference this match's nodes."""
101 return self._match_relations("referenced_by", patterns, recursive)
103 def match_references(self, patterns: Iterable[list], recursive: bool = True) -> Sequence[Self]:
104 """AI: Return matches of patterns against the nodes that this match's nodes reference."""
105 return self._match_relations("references", patterns, recursive)
107 def _match_relations(self, attr: str, patterns, recursive: bool) -> list:
108 return [
109 m
110 for node in self.nodes
111 for ref in getattr(node, attr)
112 for pattern in patterns
113 for m in MatchFinder.match_pattern([ref.node], pattern, recursive)
114 ]
116 def offset_of(self, key):
117 """AI: Return the source offset of the first node bound to the given expansion key."""
118 return self.expansions[key][0].offset
120 def length_of(self, key):
121 """AI: Return the total source span length covered by the nodes bound to the given expansion key."""
122 return self.expansions[key][-1].offset + self.expansions[key][-1].length - self.expansions[key][0].offset
125def _resolve_match_one(name: str, src: NodeProtocol, expansions: dict):
126 """Handle a MATCH_ONE pattern node: bind or verify the named expansion. Returns True if matched."""
127 if name in expansions:
128 return src == expansions[name][0]
129 expansions[name] = [src]
130 return True
133def is_match_tree(src: Sequence | None, cmp: Sequence | None, expansions=None):
134 """AI: Return True if the entire src sequence matches the cmp pattern sequence."""
135 return find_in_list(src, cmp, expansions, 0) == len(src) - 1
138def variant_in_match_stmt(src: NodeProtocol, cmp: NodeProtocol, expansions) -> list:
139 """AI: Return the list of matching Variants of a single src node against a single cmp pattern node."""
140 if pattern_kind(cmp) is PatternKind.MATCH_ONE and cmp.name:
141 matched = _resolve_match_one(cmp.name, src, expansions)
142 return [Variant(0, expansions, None, 0, 0)] if matched else []
143 if is_match_dict(src.properties, cmp.properties, expansions) and node_kinds_match(src, cmp):
144 if not cmp.children and src.children:
145 return []
146 variants = find_variants(src.children, cmp.children, expansions)
147 return [v for v in variants if v.end_index == len(src.children) - 1]
148 return []
151def _advance_match_all(variant: Variant, cmp: Sequence, src: Sequence, i: int, new_variants: list):
152 """Advance variant.index past consecutive MATCH_ALL pattern nodes, forking new_variants as needed."""
153 while pattern_kind(cmp[variant.index]) is PatternKind.MATCH_ALL:
154 current_name = cmp[variant.index].name
155 if variant.expansion_start == -1:
156 variant.expansion_start = i
157 variant.greedy = current_name
158 elif current_name != variant.greedy and variant.greedy not in variant.exp:
159 new_variants.append(variant.fork())
160 variant.close_greedy(variant.greedy, src, variant.expansion_start, i)
161 variant.greedy = current_name
162 variant.expansion_start = i
163 else:
164 break
165 has_next = (variant.index + 1) < len(cmp)
166 not_yet_expanded = current_name not in variant.exp or variant.exp[current_name] == []
167 if has_next and not_yet_expanded:
168 variant.index += 1
169 else:
170 break
173def _apply_child_match(variant: Variant, child_variants: list, cmp: Sequence, src: Sequence, i: int, new_variants: list):
174 """Apply a successful child match, forking if there are multiple child variants."""
175 greedy_open = variant.greedy is not None and variant.expansion_start != -1 and variant.greedy not in variant.exp
176 if greedy_open:
177 forked = variant.fork()
178 forked.exp.pop(cmp[variant.index].name, None)
179 new_variants.append(forked)
180 variant.close_greedy(variant.greedy, src, variant.expansion_start, i)
181 if len(child_variants) > 1:
182 new_variants.extend(Variant(variant.index + 1, v.exp, variant.greedy, variant.expansion_start) for v in child_variants)
183 variant.end_index = MIS_MATCH
184 else:
185 variant.exp = child_variants[0].exp
186 variant.index += 1
187 reached_end = i == len(src) - 1 and variant.index == len(cmp)
188 if reached_end:
189 variant.end_index = len(src) - 1
192def _advance_greedy(variant: Variant, cmp: Sequence, src: Sequence, i: int):
193 """Accumulate or verify greedy expansion for the current source node."""
194 exp_for_key = variant.exp.get(cmp[variant.index].name)
195 exp_index = i - variant.expansion_start
196 if exp_for_key is None:
197 at_last_src_node = i == len(src) - 1
198 greedy_matches_pattern = variant.greedy == cmp[variant.index].name
199 if at_last_src_node and greedy_matches_pattern:
200 variant.close_greedy(variant.greedy, src, variant.expansion_start, i + 1)
201 variant.end_index = i
202 variant.index += 1
203 elif exp_index < len(exp_for_key):
204 src_node_matches = src[i] == exp_for_key[exp_index]
205 expansion_complete = exp_index == len(exp_for_key) - 1
206 if not src_node_matches:
207 variant.end_index = MIS_MATCH
208 elif expansion_complete:
209 variant.reset_greedy()
210 variant.index += 1
211 else:
212 variant.reset_greedy()
213 variant.index += 1
216def find_variants(src: Sequence, cmp: Sequence, expansion=None, start: int = 0, parent=None):
217 """AI: Compute all candidate Variant matches of the cmp pattern sequence against the src node sequence."""
218 if expansion is None:
219 expansion = {}
220 if cmp is None:
221 return []
222 i = start
223 variants = [Variant(0, expansion, None, -1)]
224 while i < len(src):
225 next_variants = []
226 for variant in variants:
227 if variant.end_index is not INCOMPLETE_MATCH:
228 next_variants.append(variant)
229 continue
230 if variant.index == len(cmp):
231 variant.end_index = i - 1
232 next_variants.append(variant)
233 continue
234 _advance_match_all(variant, cmp, src, i, next_variants)
235 if variant.index == len(cmp):
236 next_variants.append(variant)
237 continue
238 if pattern_kind(cmp[variant.index]) is not PatternKind.MATCH_ALL and (
239 child_variants := variant_in_match_stmt(src[i], cmp[variant.index], variant.exp)
240 ):
241 _apply_child_match(variant, child_variants, cmp, src, i, next_variants)
242 elif variant.greedy:
243 _advance_greedy(variant, cmp, src, i)
244 else:
245 variant.end_index = MIS_MATCH
246 if variant.end_index != MIS_MATCH:
247 next_variants.append(variant)
248 variants = next_variants
249 i += 1
251 full_match = len(src) - 1
252 valid_variants = []
253 for variant in variants:
254 if variant.end_index == MIS_MATCH or variant.index < len(cmp) - 1:
255 continue
256 if variant.index == len(cmp) - 1:
257 last_cmp = cmp[variant.index]
258 trailing_wildcard = pattern_kind(last_cmp) is PatternKind.MATCH_ALL and last_cmp.name not in variant.exp
259 if not trailing_wildcard:
260 continue
261 key = variant.greedy if variant.expansion_start != -1 else last_cmp.name
262 variant.close_greedy(key, src, variant.expansion_start, -1 if variant.expansion_start != -1 else variant.expansion_start)
263 elif variant.index == len(cmp):
264 greedy_unresolved = variant.greedy and variant.greedy not in variant.exp
265 if greedy_unresolved:
266 variant.close_greedy(variant.greedy, src, variant.expansion_start, -1)
267 if variant.end_index == INCOMPLETE_MATCH:
268 variant.end_index = full_match
269 valid_variants.append(variant)
270 return valid_variants
273def find_in_list(src: Sequence, cmp: Sequence, exp=None, start: int = 0):
274 """AI: Return the end index of the first full match of cmp within src starting at start, or -2 if none."""
275 if exp is None:
276 exp = {}
277 variants = find_variants(src, cmp, exp, start)
278 if not variants:
279 return -2
280 exp.update(variants[0].exp)
281 # [0] most greedy
282 # [-1] least greedy
283 return variants[0].end_index
286def is_match(src: NodeProtocol, cmp: NodeProtocol, expansions=None) -> bool:
287 """AI: Return True if the single src node matches the single cmp pattern node."""
288 return variant_in_match_stmt(src, cmp, expansions) != []
291def is_match_dict(src: dict, cmp: dict, expansions: dict | None = None) -> bool:
292 """AI: Return True if all cmp properties match the corresponding src properties, binding $-placeholders."""
293 if expansions is None:
294 expansions = {}
296 def match_property(n):
297 c = cmp.get(n)
298 s = src.get(n)
299 if isinstance(c, str) and (key := use_dollar(c)).startswith("$"):
300 return s == expansions[key][0] if key in expansions else (expansions.update({key: [s]}) or True)
301 return s == c
303 return all(match_property(n) for n in (src.keys() | cmp.keys()) - IRRELEVANT_PROPS)
306def match_pattern(src_nodes, patterns, recursive=True) -> Sequence[PatternMatch]:
307 """AI: Find all matches of patterns within src_nodes, recursing into children when recursive is True."""
308 found_statements = []
309 to_do = 0
310 while to_do < len(src_nodes):
311 found_expansions = {}
312 found_position = find_in_list(src_nodes, patterns, found_expansions, to_do)
313 if found_position >= 0:
314 found_statements.append(PatternMatch(src_nodes[to_do : found_position + 1], found_expansions, patterns))
315 to_do = found_position + 1
316 else:
317 if recursive:
318 found_statements.extend(match_pattern(getattr(src_nodes[to_do], "children", []), patterns, recursive))
319 to_do += 1
320 return found_statements
323def find_all(src_nodes, *patterns, recursive: bool = True) -> Sequence[PatternMatch]:
324 """AI: Find all matches of any of the given patterns within src_nodes."""
325 return [m for pattern in patterns for m in match_pattern(src_nodes, pattern, recursive)]
328class MatchFinder:
329 """AI: Static entry points for finding pattern matches against source AST nodes."""
331 @staticmethod
332 def find_all(
333 src_nodes: Sequence[NodeProtocol],
334 *patterns: Sequence[NodeProtocol],
335 recursive: bool = True,
336 ) -> Sequence[PatternMatch]:
337 """Find all pattern matches in the given source nodes."""
338 return find_all(src_nodes, *patterns, recursive=recursive)
340 @staticmethod
341 def match_pattern(
342 src_nodes: Sequence[NodeProtocol],
343 patterns: Sequence[NodeProtocol],
344 recursive: bool = True,
345 ) -> Sequence[PatternMatch]:
346 """Match source nodes against a list of pattern nodes, optionally recursing into children."""
347 return match_pattern(src_nodes, patterns, recursive)
350# We should find the highest possible match.
351# For example in C++:
352# "int $x; $x;" matches "int x; x;"
353# "int $x = 1; int y = $x;" matches "int x = 1; int y = x;"
354# "typedef enum { $x } E; void f() { g($x); }" matches "typedef enum { x } E; void f() { g(x); }"
355# The highest shared type should be chosen as the type of $x.