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

1from collections.abc import Iterable, Sequence 

2from typing import Protocol, Self, runtime_checkable 

3 

4from renaissance.integrations.types import MatchAll, MatchOne, Type 

5from renaissance.utils.ast_utils import use_dollar 

6 

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

8 

9MIS_MATCH = -12 

10INCOMPLETE_MATCH = -11 

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

12 

13 

14@runtime_checkable 

15class AstProtocol(Protocol): 

16 ast_type: type[Type] 

17 properties: dict 

18 children: list[Self] 

19 signature: str 

20 name: str 

21 

22 

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 

30 

31 def reset_greedy(self): 

32 self.greedy = None 

33 self.expansion_start = -1 

34 

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

40 

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) 

44 

45 

46class PatternMatch: 

47 def __init__(self, nodes, expansions, patterns): 

48 self.nodes = nodes 

49 self.expansions = expansions 

50 self.patterns = patterns 

51 

52 def __str__(self): 

53 return "\n".join(node.signature for node in self.nodes) 

54 

55 @property 

56 def signature(self): 

57 return str(self) 

58 

59 def __getitem__(self, key): 

60 return "\n".join(node.signature if isinstance(node, AstProtocol) else node for node in self.expansions[key]) 

61 

62 def match_referenced_by(self, patterns: Sequence[list], recursive: bool = True) -> Sequence[Self]: 

63 return self._match_relations("referenced_by", patterns, recursive) 

64 

65 def match_references(self, patterns: Iterable[list], recursive: bool = True) -> Sequence[Self]: 

66 return self._match_relations("references", patterns, recursive) 

67 

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 ] 

76 

77 def offset_of(self, key): 

78 return self.expansions[key][0].offset 

79 

80 def length_of(self, key): 

81 return self.expansions[key][-1].offset + self.expansions[key][-1].length - self.expansions[key][0].offset 

82 

83 

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 

90 

91 

92def is_match_tree(src: Sequence | None, cmp: Sequence | None, expansions=None): 

93 return find_in_list(src, cmp, expansions, 0) == len(src) - 1 

94 

95 

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

106 

107 

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 

128 

129 

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 

148 

149 

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 

172 

173 

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 

207 

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 

228 

229 

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 

240 

241 

242def is_match(src: AstProtocol, cmp: AstProtocol, expansions=None) -> bool: 

243 return variant_in_match_stmt(src, cmp, expansions) != [] 

244 

245 

246def is_match_dict(src: dict, cmp: dict, expansions: dict | None = None) -> bool: 

247 if expansions is None: 

248 expansions = {} 

249 

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 

256 

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

258 

259 

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 

274 

275 

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

278 

279 

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) 

289 

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) 

298 

299 

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.