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

1"""AI: Pattern matching engine that finds AST nodes matching a given pattern.""" 

2 

3from collections.abc import Iterable, Sequence 

4from typing import Self 

5 

6from renaissance.utils.ast_utils import use_dollar 

7 

8from .node_protocol import NodeProtocol 

9from .pattern_kind import PatternKind 

10from .semantic_kind import SemanticKind 

11 

12IRRELEVANT_PROPS = {"macro_expansion", "start_point", "end_point", "source_code", "location", "type"} 

13 

14MIS_MATCH = -12 

15INCOMPLETE_MATCH = -11 

16_TOP_LEVEL_KINDS = {"Module", "TRANSLATION_UNIT"} 

17 

18 

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 

29 

30 

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 

41 

42 

43class Variant: 

44 """AI: Track one candidate pattern-match state (bound expansions, greedy position) during matching.""" 

45 

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 

60 

61 def reset_greedy(self): 

62 """AI: Clear the current greedy-expansion tracking state.""" 

63 self.greedy = None 

64 self.expansion_start = -1 

65 

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() 

71 

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) 

75 

76 

77class PatternMatch: 

78 """AI: Represent a successful match of a pattern against a sequence of AST nodes.""" 

79 

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 

85 

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) 

89 

90 @property 

91 def signature(self): 

92 """AI: Return the newline-joined signatures of the matched nodes.""" 

93 return str(self) 

94 

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]) 

98 

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) 

102 

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) 

106 

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 ] 

115 

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 

119 

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 

123 

124 

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 

131 

132 

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 

136 

137 

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 [] 

149 

150 

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 

171 

172 

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 

190 

191 

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 

214 

215 

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 

250 

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 

271 

272 

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 

284 

285 

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) != [] 

289 

290 

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 = {} 

295 

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 

302 

303 return all(match_property(n) for n in (src.keys() | cmp.keys()) - IRRELEVANT_PROPS) 

304 

305 

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 

321 

322 

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)] 

326 

327 

328class MatchFinder: 

329 """AI: Static entry points for finding pattern matches against source AST nodes.""" 

330 

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) 

339 

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) 

348 

349 

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.