Coverage for src / renaissance / syntax_tree / match_finder.py: 98%
210 statements
« prev ^ index » next coverage.py v7.13.4, created at 2026-09-09 14:04 +0000
« prev ^ index » next coverage.py v7.13.4, created at 2026-09-09 14:04 +0000
1from collections.abc import Iterable, Sequence
2from typing import Protocol, Self, runtime_checkable
4from renaissance.integrations.types import MatchAll, MatchOne, Type
5from renaissance.utils.ast_utils import use_dollar
7IRRELEVANT_PROPS = {"macro_expansion", "start_point", "end_point", "source_code", "location", "type"}
9MIS_MATCH = -12
10INCOMPLETE_MATCH = -11
11_TOP_LEVEL_KINDS = {"Module", "TRANSLATION_UNIT"}
14@runtime_checkable
15class AstProtocol(Protocol):
16 ast_type: type[Type]
17 properties: dict
18 children: list[Self]
19 signature: str
20 name: str
23class Variant:
24 def __init__(self, index, exp, greedy, expansion_start, end_index=INCOMPLETE_MATCH):
25 self.exp: dict = exp
26 self.index: int = index
27 self.greedy: str = greedy
28 self.end_index = end_index
29 self.expansion_start = expansion_start
31 def reset_greedy(self):
32 self.greedy = None
33 self.expansion_start = -1
35 def close_greedy(self, key, nodes, start, end):
36 """Store a completed greedy expansion and reset greedy state."""
37 value = nodes[start:end]
38 self.exp[key] = value
39 self.reset_greedy()
41 def fork(self) -> Variant:
42 """Return a copy of this variant at the same position (for backtracking)."""
43 return Variant(self.index, self.exp.copy(), self.greedy, self.expansion_start)
46class PatternMatch:
47 def __init__(self, nodes, expansions, patterns):
48 self.nodes = nodes
49 self.expansions = expansions
50 self.patterns = patterns
52 def __str__(self):
53 return "\n".join(node.signature for node in self.nodes)
55 @property
56 def signature(self):
57 return str(self)
59 def __getitem__(self, key):
60 return "\n".join(node.signature if isinstance(node, AstProtocol) else node for node in self.expansions[key])
62 def match_referenced_by(self, patterns: Sequence[list], recursive: bool = True) -> Sequence[Self]:
63 return self._match_relations("referenced_by", patterns, recursive)
65 def match_references(self, patterns: Iterable[list], recursive: bool = True) -> Sequence[Self]:
66 return self._match_relations("references", patterns, recursive)
68 def _match_relations(self, attr: str, patterns, recursive: bool) -> list:
69 return [
70 m
71 for node in self.nodes
72 for ref in getattr(node, attr)
73 for pattern in patterns
74 for m in MatchFinder.match_pattern([ref.node], pattern, recursive)
75 ]
77 def offset_of(self, key):
78 return self.expansions[key][0].offset
80 def length_of(self, key):
81 return self.expansions[key][-1].offset + self.expansions[key][-1].length - self.expansions[key][0].offset
84def _resolve_match_one(name: str, src: AstProtocol, expansions: dict):
85 """Handle a MATCH_ONE pattern node: bind or verify the named expansion. Returns True if matched."""
86 if name in expansions:
87 return src == expansions[name][0]
88 expansions[name] = [src]
89 return True
92def is_match_tree(src: Sequence | None, cmp: Sequence | None, expansions=None):
93 return find_in_list(src, cmp, expansions, 0) == len(src) - 1
96def variant_in_match_stmt(src: AstProtocol, cmp: AstProtocol, expansions) -> list:
97 if cmp.ast_type == MatchOne and cmp.name:
98 matched = _resolve_match_one(cmp.name, src, expansions)
99 return [Variant(0, expansions, None, 0, 0)] if matched else []
100 if is_match_dict(src.properties, cmp.properties, expansions) and src.ast_type == cmp.ast_type:
101 if not cmp.children and src.children:
102 return []
103 variants = find_variants(src.children, cmp.children, expansions)
104 return [v for v in variants if v.end_index == len(src.children) - 1]
105 return []
108def _advance_match_all(variant: Variant, cmp: Sequence, src: Sequence, i: int, new_variants: list):
109 """Advance variant.index past consecutive MATCH_ALL pattern nodes, forking new_variants as needed."""
110 while cmp[variant.index].ast_type == MatchAll:
111 current_name = cmp[variant.index].name
112 if variant.expansion_start == -1:
113 variant.expansion_start = i
114 variant.greedy = current_name
115 elif current_name != variant.greedy and variant.greedy not in variant.exp:
116 new_variants.append(variant.fork())
117 variant.close_greedy(variant.greedy, src, variant.expansion_start, i)
118 variant.greedy = current_name
119 variant.expansion_start = i
120 else:
121 break
122 has_next = (variant.index + 1) < len(cmp)
123 not_yet_expanded = current_name not in variant.exp or variant.exp[current_name] == []
124 if has_next and not_yet_expanded:
125 variant.index += 1
126 else:
127 break
130def _apply_child_match(variant: Variant, child_variants: list, cmp: Sequence, src: Sequence, i: int, new_variants: list):
131 """Apply a successful child match, forking if there are multiple child variants."""
132 greedy_open = variant.greedy is not None and variant.expansion_start != -1 and variant.greedy not in variant.exp
133 if greedy_open:
134 forked = variant.fork()
135 forked.exp.pop(cmp[variant.index].name, None)
136 new_variants.append(forked)
137 variant.close_greedy(variant.greedy, src, variant.expansion_start, i)
138 if len(child_variants) > 1:
139 for v in child_variants:
140 new_variants.append(Variant(variant.index + 1, v.exp, variant.greedy, variant.expansion_start))
141 variant.end_index = MIS_MATCH
142 else:
143 variant.exp = child_variants[0].exp
144 variant.index += 1
145 reached_end = i == len(src) - 1 and variant.index == len(cmp)
146 if reached_end:
147 variant.end_index = len(src) - 1
150def _advance_greedy(variant: Variant, cmp: Sequence, src: Sequence, i: int):
151 """Accumulate or verify greedy expansion for the current source node."""
152 exp_for_key = variant.exp.get(cmp[variant.index].name)
153 exp_index = i - variant.expansion_start
154 if exp_for_key is None:
155 at_last_src_node = i == len(src) - 1
156 greedy_matches_pattern = variant.greedy == cmp[variant.index].name
157 if at_last_src_node and greedy_matches_pattern:
158 variant.close_greedy(variant.greedy, src, variant.expansion_start, i + 1)
159 variant.end_index = i
160 variant.index += 1
161 elif exp_index < len(exp_for_key):
162 src_node_matches = src[i] == exp_for_key[exp_index]
163 expansion_complete = exp_index == len(exp_for_key) - 1
164 if not src_node_matches:
165 variant.end_index = MIS_MATCH
166 elif expansion_complete:
167 variant.reset_greedy()
168 variant.index += 1
169 else:
170 variant.reset_greedy()
171 variant.index += 1
174def find_variants(src: Sequence, cmp: Sequence, expansion=None, start: int = 0, parent=None):
175 if expansion is None:
176 expansion = {}
177 if cmp is None:
178 return []
179 i = start
180 variants = [Variant(0, expansion, None, -1)]
181 while i < len(src):
182 next_variants = []
183 for variant in variants:
184 if variant.end_index is not INCOMPLETE_MATCH:
185 next_variants.append(variant)
186 continue
187 if variant.index == len(cmp):
188 variant.end_index = i - 1
189 next_variants.append(variant)
190 continue
191 _advance_match_all(variant, cmp, src, i, next_variants)
192 if variant.index == len(cmp):
193 next_variants.append(variant)
194 continue
195 if cmp[variant.index].ast_type != MatchAll and (
196 child_variants := variant_in_match_stmt(src[i], cmp[variant.index], variant.exp)
197 ):
198 _apply_child_match(variant, child_variants, cmp, src, i, next_variants)
199 elif variant.greedy:
200 _advance_greedy(variant, cmp, src, i)
201 else:
202 variant.end_index = MIS_MATCH
203 if variant.end_index != MIS_MATCH:
204 next_variants.append(variant)
205 variants = next_variants
206 i += 1
208 full_match = len(src) - 1
209 valid_variants = []
210 for variant in variants:
211 if variant.end_index == MIS_MATCH or variant.index < len(cmp) - 1:
212 continue
213 if variant.index == len(cmp) - 1:
214 last_cmp = cmp[variant.index]
215 trailing_wildcard = last_cmp.ast_type == MatchAll and last_cmp.name not in variant.exp
216 if not trailing_wildcard:
217 continue
218 key = variant.greedy if variant.expansion_start != -1 else last_cmp.name
219 variant.close_greedy(key, src, variant.expansion_start, -1 if variant.expansion_start != -1 else variant.expansion_start)
220 elif variant.index == len(cmp):
221 greedy_unresolved = variant.greedy and variant.greedy not in variant.exp
222 if greedy_unresolved:
223 variant.close_greedy(variant.greedy, src, variant.expansion_start, -1)
224 if variant.end_index == INCOMPLETE_MATCH:
225 variant.end_index = full_match
226 valid_variants.append(variant)
227 return valid_variants
230def find_in_list(src: Sequence, cmp: Sequence, exp=None, start: int = 0):
231 if exp is None:
232 exp = {}
233 variants = find_variants(src, cmp, exp, start)
234 if not variants:
235 return -2
236 exp.update(variants[0].exp)
237 # [0] most greedy
238 # [-1] least greedy
239 return variants[0].end_index
242def is_match(src: AstProtocol, cmp: AstProtocol, expansions=None) -> bool:
243 return variant_in_match_stmt(src, cmp, expansions) != []
246def is_match_dict(src: dict, cmp: dict, expansions: dict | None = None) -> bool:
247 if expansions is None:
248 expansions = {}
250 def match_property(n):
251 c = cmp.get(n)
252 s = src.get(n)
253 if isinstance(c, str) and (key := use_dollar(c)).startswith("$"):
254 return s == expansions[key][0] if key in expansions else (expansions.update({key: [s]}) or True)
255 return s == c
257 return all(match_property(n) for n in (src.keys() | cmp.keys()) - IRRELEVANT_PROPS)
260def match_pattern(src_nodes, patterns, recursive=True) -> Sequence[PatternMatch]:
261 found_statements = []
262 to_do = 0
263 while to_do < len(src_nodes):
264 found_expansions = {}
265 found_position = find_in_list(src_nodes, patterns, found_expansions, to_do)
266 if found_position >= 0:
267 found_statements.append(PatternMatch(src_nodes[to_do : found_position + 1], found_expansions, patterns))
268 to_do = found_position + 1
269 else:
270 if recursive:
271 found_statements.extend(match_pattern(getattr(src_nodes[to_do], "children", []), patterns, recursive))
272 to_do += 1
273 return found_statements
276def find_all(src_nodes, *patterns, recursive: bool = True) -> Sequence[PatternMatch]:
277 return [m for pattern in patterns for m in match_pattern(src_nodes, pattern, recursive)]
280class MatchFinder:
281 @staticmethod
282 def find_all(
283 src_nodes: Sequence[AstProtocol],
284 *patterns: Sequence[AstProtocol],
285 recursive: bool = True,
286 ) -> Sequence[PatternMatch]:
287 """Finds all pattern matches in the given source nodes."""
288 return find_all(src_nodes, *patterns, recursive=recursive)
290 @staticmethod
291 def match_pattern(
292 src_nodes: Sequence[AstProtocol],
293 patterns: Sequence[AstProtocol],
294 recursive: bool = True,
295 ) -> Sequence[PatternMatch]:
296 """Matches source nodes against a list of pattern nodes, optionally recursing into children."""
297 return match_pattern(src_nodes, patterns, recursive)
300# We should find the highest possible match.
301# For example in C++:
302# "int $x; $x;" matches "int x; x;"
303# "int $x = 1; int y = $x;" matches "int x = 1; int y = x;"
304# "typedef enum { $x } E; void f() { g($x); }" matches "typedef enum { x } E; void f() { g(x); }"
305# The highest shared type should be chosen as the type of $x.