Coverage for src/renaissance/integrations/python/ast/rst_node.py: 96%

332 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-10-02 13:12 +0000

1"""AI: ASTNode implementation backed by Python's built-in `ast` module (the "RST" representation).""" 

2 

3import ast 

4import sys 

5import textwrap 

6from collections.abc import Callable, Sequence 

7from pathlib import Path 

8from typing import Any, Self 

9 

10from renaissance.integrations.python.ast.kinds import PYTHON_KIND_MAP, PYTHON_OPERATOR_MAP 

11from renaissance.integrations.python.ast.util import convert 

12from renaissance.syntax_tree.match_finder import find_in_list 

13from renaissance.syntax_tree.semantic_kind import SemanticKind 

14from renaissance.utils.ast_utils import ( 

15 format_node, 

16 match_children, 

17 match_props, 

18 next_sibling, 

19 preceding_sibling, 

20 traverse, 

21) 

22 

23types = ["int", "float", "str", "list", "set", "tuple", "Mapping", "dict", "Optional"] 

24IRRELEVANT_PROPS = {"comment"} 

25IRRELEVANT_NODES = {"comment"} 

26IMPLICIT = {"ImplicitNode"} 

27 

28 

29class ImplicitNode(ast.Name): 

30 """AI: Represent a synthetic AST node inserted where the real source has none.""" 

31 

32 _fields = ( 

33 "id", 

34 "body", 

35 ) 

36 

37 _field_types = { 

38 "id": str, 

39 "body": list, 

40 } 

41 

42 def __init__(self, name, children=None): 

43 """AI: Represent a synthetic AST node inserted where the real source has none.""" 

44 super().__init__(name, children or []) 

45 self.lineno = 0 

46 self.col_offset = 0 

47 self.end_lineno = 0 

48 self.end_col_offset = 0 

49 

50 

51class PythonRSTReference: 

52 """AI: Represent a reference from one Python RST AST node to another by id.""" 

53 

54 def __repr__(self): 

55 """AI: Return a string identifying the referenced node id and reference kind.""" 

56 return f"{self.node_id}:{self.ref_kind}" 

57 

58 def __init__(self, node_id: str, ref_kind: str, properties: dict[str, Any]) -> None: 

59 """AI: Represent a reference from one Python RST AST node to another by id.""" 

60 self.node_id = node_id 

61 self.ref_kind = ref_kind 

62 self.properties = properties 

63 

64 

65class PythonRstTranslationUnit: 

66 """AI: Parse Python source into a stdlib ast tree with lazily-built reference caches.""" 

67 

68 cache = {} 

69 

70 def __init__(self, content, file_name: str): 

71 """AI: Parse Python source into a stdlib ast tree with lazily-built reference caches.""" 

72 self.content = content.encode(sys.getfilesystemencoding()) 

73 self.atu = ast.parse(content, file_name) 

74 self.file_name = file_name 

75 self.references_initialized = False 

76 PythonRstTranslationUnit.cache[file_name] = content 

77 self.lines = self.content.splitlines() 

78 

79 self._references: dict[str, list[PythonRSTReference]] = {} 

80 self._referenced_by: dict[str, list[PythonRSTReference]] = {} 

81 self._nodes: dict[str, PythonRstNode] = {} 

82 

83 def check_diagnostics(self, continue_with_warning=True) -> None: 

84 """AI: Report ast type-ignore comments, raising an Exception unless continue_with_warning is True.""" 

85 msg = None 

86 errors = "" 

87 for d in self.atu.type_ignores: 

88 msg = f"type ignored: {d.tag} at {d.lineno}\n" 

89 errors += msg 

90 print(msg) 

91 if msg and not continue_with_warning: 

92 raise Exception(f"Error parsing: {self.file_name} \n+ errors: {errors}") 

93 

94 def lazy_create_refers(self, node: PythonRstNode) -> None: 

95 """AI: Build the translation unit's reference cache on first use, then no-op on subsequent calls.""" 

96 if self.references_initialized: 

97 return 

98 for n in traverse(node.root): 

99 self.create_references(n) 

100 self.references_initialized = True 

101 

102 def add(self, node): 

103 """AI: Register node in the node-name lookup table, keyed by its parser-kind-specific name.""" 

104 match node.parser_kind: 

105 case "Name": 

106 if node.node.id not in self._nodes and node.node.id not in types: 

107 self._nodes[node.node.id] = node 

108 case "FunctionDef": 

109 if node.node.name not in self._nodes: 

110 self._nodes[node.node.name] = node 

111 case "Call": 

112 if node.name not in self._nodes: 

113 self._nodes[node.name] = node 

114 case "ClassDef": 

115 if node.name not in self._nodes: 

116 self._nodes[node.name] = node 

117 case "arg": 

118 if node.name != "self" and node.name not in self._nodes: 

119 self._nodes[node.name] = node 

120 

121 def create_references(self, ast_node) -> None: 

122 """AI: Derive and record type, call, and inheritance references for ast_node based on its ast type.""" 

123 assert isinstance(ast_node, PythonRstNode), f"Expected PythonASTNode but got {type(ast_node)}" 

124 match type(ast_node.node): 

125 case ast.arg: 

126 # TODO: is excluding self here intentional..? 

127 if ast_node.name != "self" and isinstance(ast_node.node, ast.arg) and isinstance(ast_node.node.annotation, ast.Name): 

128 node_id = ast_node.name 

129 ref_id = ast_node.node.annotation.id 

130 ref_kind = "TypeRef" 

131 self.add_reference(node_id, ref_id, ref_kind) 

132 case ast.FunctionDef | ast.AsyncFunctionDef: 

133 if isinstance(ast_node.node, (ast.FunctionDef, ast.AsyncFunctionDef)) and isinstance(ast_node.node.returns, ast.Name): 

134 node_id = ast_node.name 

135 ref_id = ast_node.node.returns.id 

136 ref_kind = "TypeRef" 

137 self.add_reference(node_id, ref_id, ref_kind) 

138 case ast.Assign: 

139 if isinstance(ast_node.node, ast.Assign): 

140 for n in ast_node.node.targets: 

141 if isinstance(n, ast.Name) and isinstance(ast_node.node.value, ast.Call): 

142 node_id = n.id 

143 func = ast_node.node.value.func 

144 ref_id = func.id if isinstance(func, ast.Name) else None 

145 if ref_id: 

146 ref_kind = "CallRef" 

147 self.add_reference(node_id, ref_id, ref_kind) 

148 case ast.AnnAssign: 

149 if isinstance(ast_node.node, ast.AnnAssign) and ( 

150 ast_node.node.annotation 

151 and isinstance(ast_node.node.target, ast.Name) 

152 and isinstance(ast_node.node.annotation, ast.Name) 

153 ): 

154 node_id = ast_node.node.target.id 

155 ref_id = ast_node.node.annotation.id 

156 ref_kind = "TypeRef" 

157 self.add_reference(node_id, ref_id, ref_kind) 

158 case ast.ClassDef: 

159 if isinstance(ast_node.node, ast.ClassDef): 

160 node = ast_node.node 

161 node_id = node.name 

162 if node.bases: 

163 ref_node = node.bases[0] 

164 if isinstance(ref_node, ast.Name): 

165 ref_id = ref_node.id 

166 ref_kind = "Inherit" 

167 self.add_reference(node_id, ref_id, ref_kind) 

168 # add functions and attributes to class 

169 

170 case ast.Call: 

171 if isinstance(ast_node.node, ast.Call): 

172 # obj.function. then obj refers to function 

173 if isinstance(ast_node.node.func, ast.Attribute): 

174 node_id = ast_node.name 

175 ref_id = ast_node.node.func.attr 

176 ref_kind = "FuncCall" 

177 self.add_reference(node_id, ref_id, ref_kind) 

178 # call function 'a' in function 'b', then 'b' refers to 'a' 

179 container = ast_node.get_container_parent() 

180 if container.parser_kind in {"FunctionDef", "AsyncFunctionDef"} and isinstance(ast_node.node.func, ast.Name): 

181 node_id = container.name 

182 ref_id = ast_node.node.func.id 

183 ref_kind = "FuncCall" 

184 self.add_reference(node_id, ref_id, ref_kind) 

185 

186 def add_reference(self, node_id: str, ref_id: str, ref_kind: str) -> None: 

187 """AI: Record a reference from node_id to ref_id (and the corresponding referenced-by entry).""" 

188 properties = {} 

189 if node_id == ref_id: 

190 return 

191 reference = PythonRSTReference(ref_id, ref_kind, properties) 

192 referenced_by = PythonRSTReference(node_id, ref_kind, properties) 

193 if node_id in self._references: 

194 self._references[node_id].append(reference) 

195 else: 

196 self._references[node_id] = [reference] 

197 if ref_id in self._referenced_by: 

198 self._referenced_by[ref_id].append(referenced_by) 

199 else: 

200 self._referenced_by[ref_id] = [referenced_by] 

201 

202 def get_referenced_by(self, node_id): 

203 """AI: Return the references pointing to node_id, resolved to their referring nodes' names.""" 

204 refs = self._referenced_by.get(node_id, []) 

205 return [PythonRSTReference(self._nodes[ref.node_id].name, ref.ref_kind, ref.properties) for ref in refs] 

206 

207 def get_references(self, node_id): 

208 """AI: Return the references that node_id points to, resolved to their target nodes' names.""" 

209 refs = self._references.get(node_id, []) 

210 return [PythonRSTReference(self._nodes[ref.node_id].name, ref.ref_kind, ref.properties) for ref in refs] 

211 

212 

213class PythonRstNode: 

214 """AI: ASTNode implementation backed by the stdlib ast module.""" 

215 

216 def __init__(self, node: ast.AST, translation_unit: PythonRstTranslationUnit = None, parent=None): 

217 """AI: Wrap a stdlib ast node as an AST node within the given translation unit.""" 

218 self.root = parent.root if parent and parent.root else self 

219 self.node = node 

220 self.parent = parent 

221 self.translation_unit: PythonRstTranslationUnit = translation_unit 

222 self.parser_kind = type(node).__name__ 

223 self.semantic_kind = PYTHON_KIND_MAP.get(self.parser_kind, SemanticKind.NODE) 

224 

225 self.indent = "" 

226 self.name = self._derive_name() 

227 self.show_props = False 

228 self.children = [] 

229 self.properties = {} 

230 self.is_implicit = self.parser_kind not in IMPLICIT 

231 self.offset = 0 

232 self.length = 0 

233 if self.translation_unit: 

234 self.filename = translation_unit.file_name 

235 self.derive_position(node, translation_unit, parent) 

236 self.add_node() 

237 for name in node._fields: 

238 try: 

239 child = getattr(node, name) 

240 match child: 

241 case list(): # Matches any list 

242 if isinstance(node, ast.Global) and name == "names" and len(child) == 1: 

243 self.name = child[0] 

244 

245 if isinstance(node, (ast.Global, ast.Nonlocal)): 

246 pass # names is list[str], not AST nodes - nothing to build children from 

247 elif isinstance(node, (ImplicitNode, ast.Module)) or len(node._fields) == 1: 

248 # A list field can hold a bare None at a position with no value 

249 # None isn't a real AST node, so it has nothing to build a child from. 

250 self.children.extend(PythonRstNode(n, translation_unit, self) for n in child if n is not None) 

251 if name == "body": 

252 self.body = self.children 

253 else: 

254 self.children.append(PythonRstNode(ImplicitNode(name, child), translation_unit, self)) 

255 if name in ["body", "cases"]: 

256 self.body = self.children[-1].children 

257 

258 case ast.AST(): 

259 if name != "ctx": 

260 self.children.append(PythonRstNode(child, translation_unit, self)) 

261 if isinstance(child, ast.expr): 

262 self.expression = self.children[-1] 

263 case _: 

264 if name != "None": 

265 self.properties[name] = child 

266 except AttributeError as e: 

267 print(e) 

268 continue 

269 

270 self.end_offset = self.offset + self.length 

271 self.extended_end_offset = self.end_offset 

272 self.is_statement = isinstance(self.node, ast.stmt) 

273 

274 @property 

275 def kind_key(self) -> SemanticKind | str: 

276 """AI: Return the semantic kind, or the raw parser kind when no semantic kind applies.""" 

277 return self.semantic_kind if self.semantic_kind is not SemanticKind.NODE else self.parser_kind 

278 

279 def __eq__(self, other): 

280 """AI: Return whether this node is structurally equal to `other`, ignoring irrelevant properties/children.""" 

281 return ( 

282 isinstance(other, type(self)) 

283 and self.kind_key == other.kind_key 

284 and match_props(self.properties, other.properties, IRRELEVANT_PROPS) 

285 and match_children(self.children, other.children, IRRELEVANT_NODES) 

286 ) 

287 

288 def __contains__(self, item): 

289 """AI: Return whether item(s) are found among this node's children.""" 

290 if not isinstance(item, list): 

291 item = [item] 

292 return find_in_list(self.children, item) 

293 

294 def __getitem__(self, key): 

295 """Allow indexing/slicing into node to access children. 

296 

297 Usage: node[0] == node.children[0] 

298 """ 

299 return self.children[key] 

300 

301 def __repr__(self): 

302 """AI: Return the formatted node representation.""" 

303 return format_node(self) 

304 

305 @property 

306 def next_sibling(self) -> Self | None: 

307 """AI: Return the sibling node immediately following this one, or None.""" 

308 return next_sibling(self) 

309 

310 @property 

311 def preceding_sibling(self) -> Self | None: 

312 """AI: Return the sibling node immediately preceding this one, or None.""" 

313 return preceding_sibling(self) 

314 

315 def process(self, function: Callable[[Self], None]) -> None: 

316 """AI: Apply function to this node and recursively to all of its descendants.""" 

317 function(self) 

318 for child in self.children: 

319 child.process(function) 

320 

321 def derive_position(self, node: ast.AST, translation_unit: PythonRstTranslationUnit, parent): 

322 """AI: Compute and set this node's offset and length in the source text from its ast position attributes.""" 

323 if node._attributes: 

324 if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)) and node.decorator_list: 

325 self.offset = convert(self.translation_unit.lines, node.decorator_list[0].lineno, node.decorator_list[0].col_offset) - 1 

326 elif parent.name == "decorator_list": 

327 # also include the @ in the decorator 

328 self.offset = convert(self.translation_unit.lines, node.lineno, node.col_offset) - 1 # type: ignore[attr-defined] 

329 else: 

330 self.offset = convert(self.translation_unit.lines, node.lineno, node.col_offset) # type: ignore[attr-defined] 

331 all_space = all(c == " " for c in self.translation_unit.content[self.offset - node.col_offset : self.offset]) 

332 if all_space: 

333 self.offset = max(self.offset - node.col_offset, 0) 

334 self.length = convert(self.translation_unit.lines, node.end_lineno, node.end_col_offset) - self.offset # type: ignore[attr-defined] 

335 elif isinstance(node, ast.Module) and translation_unit: 

336 self.offset = 0 

337 self.length = len(translation_unit.content) 

338 else: 

339 self.offset = 0 

340 self.length = 0 

341 

342 @staticmethod 

343 def load( 

344 file_path: Path, # TODO: why Path - why not FileDescriptorOrPath (the type of the file parameter of the open function)? 

345 extra_args: Sequence[str] | None = None, 

346 working_dir: Path | None = None, 

347 ) -> PythonRstNode: 

348 """AI: Parse the Python source file at file_path into a PythonRstNode tree.""" 

349 # Keep a uniform loader signature across AST node implementations. 

350 # Python's AST parser does not need extra arguments or a working dir. 

351 _ = extra_args, working_dir 

352 with Path(file_path).open() as file: 

353 content = file.read() 

354 return PythonRstNode.load_from_text(content, str(file_path)) 

355 

356 @staticmethod 

357 def load_from_text( 

358 text: str, 

359 file_name: str = "test.py", 

360 extra_args: Sequence[str] | None = None, 

361 working_dir: Path | None = None, 

362 ) -> PythonRstNode: 

363 """AI: Parse Python source text into a PythonRstNode tree, attributed to file_name.""" 

364 _ = extra_args, working_dir 

365 translation_unit = PythonRstTranslationUnit(text, file_name=str(file_name)) 

366 translation_unit.check_diagnostics() 

367 root_node = PythonRstNode(translation_unit.atu, translation_unit) 

368 return root_node 

369 

370 def _derive_name(self): 

371 

372 if ( 

373 isinstance( 

374 self.node, 

375 ( 

376 ast.FunctionDef, 

377 ast.AsyncFunctionDef, 

378 ast.ClassDef, 

379 ast.ExceptHandler, 

380 ), 

381 ) 

382 and self.node.name 

383 ): 

384 name = self.node.name 

385 elif isinstance(self.node, ast.Global) and len(self.node.names) == 1: 

386 name = self.node.names[0] 

387 elif isinstance(self.node, (ast.AnnAssign, ast.AugAssign)) and isinstance(self.node.target, ast.Name): 

388 name = self.node.target.id 

389 elif isinstance(self.node, ast.Assign) and len(self.node.targets) == 1: 

390 target = self.node.targets[0] 

391 name = target.id if isinstance(target, ast.Name) else self.parser_kind 

392 elif isinstance(self.node, ast.Name): 

393 name = self.node.id 

394 elif isinstance(self.node, ast.arg): 

395 name = self.node.arg 

396 elif isinstance(self.node, ast.Match) and isinstance(self.node.subject, ast.Name): 

397 name = self.node.subject.id 

398 elif (isinstance(self.node, ast.Import) and len(self.node.names) == 1) or ( 

399 isinstance(self.node, ast.ImportFrom) and len(self.node.names) == 1 

400 ): 

401 name = self.node.names[0].name 

402 elif isinstance(self.node, (ast.Assert, ast.Break, ast.Pass, ast.Raise, ast.Continue)): 

403 name = "" 

404 elif isinstance(self.node, (ast.For, ast.AsyncFor)): 

405 if ( 

406 isinstance(self.node.target, ast.Tuple) 

407 and len(self.node.target.elts) > 1 

408 and isinstance(self.node.target.elts[1], ast.Name) 

409 ): 

410 name = self.node.target.elts[1].id 

411 elif isinstance(self.node.target, ast.Name): 

412 name = self.node.target.id 

413 else: 

414 name = str(self.node.target) 

415 elif "body" not in self.node._fields: 

416 name = ast.unparse(self.node) 

417 elif isinstance(self.node, (ast.Module)) and self.translation_unit: 

418 name = self.translation_unit.file_name 

419 else: 

420 name = self.parser_kind 

421 return name or "" 

422 

423 @property 

424 def type(self): 

425 """AI: Return the annotation name for an annotated assignment node, or None otherwise.""" 

426 return self.node.annotation.id if isinstance(self.node, ast.AnnAssign) and isinstance(self.node.annotation, ast.Name) else None 

427 

428 @property 

429 def value(self): 

430 """AI: Return the literal value of this node, or None if it has none.""" 

431 if self.parser_kind == "Assert": 

432 return 0 

433 return self.node.value.value if hasattr(self.node, "value") else None 

434 

435 @property 

436 def expr(self): 

437 """AI: Return the wrapped inner expression node relevant to this statement's kind, or None.""" 

438 if ( 

439 isinstance( 

440 self.node, 

441 ( 

442 ast.Assign, 

443 ast.AnnAssign, 

444 ast.AugAssign, 

445 ast.Return, 

446 ast.Expr, 

447 ast.Delete, 

448 ast.NamedExpr, 

449 ), 

450 ) 

451 and hasattr(self.node, "value") 

452 and self.node.value is not None 

453 ) or (isinstance(self.node, ast.Expr) and hasattr(self.node, "value")): 

454 return PythonRstNode(self.node.value, self.translation_unit, self) 

455 if isinstance(self.node, (ast.For, ast.AsyncFor, ast.comprehension)): 

456 return PythonRstNode(self.node.iter, self.translation_unit, self) 

457 if isinstance(self.node, (ast.If, ast.While, ast.Assert)): 

458 return PythonRstNode(self.node.test, self.translation_unit, self) 

459 if isinstance(self.node, (ast.Raise, ast.ExceptHandler)) and hasattr(self.node, "exc") and self.node.exc is not None: 

460 return PythonRstNode(self.node.exc, self.translation_unit, self) 

461 return None 

462 

463 @property 

464 def operator(self): 

465 """AI: Return this node's operator symbol, or an empty string if it has none.""" 

466 node_type = type(self.node).__name__ 

467 op = type(self.node.op).__name__ if isinstance(self.node, (ast.BinOp, ast.UnaryOp, ast.BoolOp, ast.AugAssign)) else "" 

468 return PYTHON_OPERATOR_MAP.get(node_type + op, "") 

469 

470 @property 

471 def signature(self) -> str: 

472 """AI: Return the source code text of this node, prefixed with '@' for decorators.""" 

473 sig = self.binary_file_content().decode(sys.getfilesystemencoding()) 

474 if self.parent and self.parent.name == "decorator_list" and not sig.startswith("@"): 

475 sig = "@" + sig 

476 return sig 

477 

478 def binary_file_content(self) -> bytes: 

479 """AI: Return this node's source text as encoded bytes.""" 

480 return ( 

481 self.translation_unit.content[self.offset : self.offset + self.length] 

482 if self.translation_unit 

483 else ast.unparse(self.node).encode(sys.getfilesystemencoding()) 

484 ) 

485 

486 @property 

487 def referenced_by(self) -> Sequence[PythonRSTReference]: 

488 """AI: Return the references that point to this node, building the reference cache if needed.""" 

489 self.translation_unit.lazy_create_refers(self) 

490 return self.translation_unit.get_referenced_by(self.name) 

491 

492 @property 

493 def references(self) -> list[PythonRSTReference]: 

494 """AI: Return the references that this node points to, building the reference cache if needed.""" 

495 self.translation_unit.lazy_create_refers(self) 

496 return self.translation_unit.get_references(self.name) 

497 

498 def add_node(self): 

499 """AI: Register this node in its translation unit's node-name lookup table.""" 

500 self.translation_unit.add(self) 

501 

502 def get_container_parent(self): 

503 """AI: Return the nearest ancestor node that is a function, class, or module.""" 

504 if self.parent: 

505 if self.parent.parser_kind in {"FunctionDef", "AsyncFunctionDef", "ClassDef", "Module"}: 

506 return self.parent 

507 return self.parent.get_container_parent() 

508 return self 

509 

510 @property 

511 def text(self) -> str: 

512 """AI: Return this node's signature text with common leading whitespace removed.""" 

513 return textwrap.dedent(self.signature)