Coverage for src/renaissance/syntax_tree/ast_processor.py: 85%

68 statements  

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

1"""AI: Processor that coordinates parsing, pattern matching, and rewriting of a single AST.""" 

2 

3from collections.abc import Callable, Iterator, Sequence 

4from pathlib import Path 

5 

6import renaissance.syntax_tree.match_finder 

7from renaissance.syntax_tree import ASTNode 

8from renaissance.syntax_tree.ast_factory import ASTFactory 

9from renaissance.syntax_tree.ast_finder import ASTFinder, find_semantic_kind 

10from renaissance.syntax_tree.ast_rewriter import ASTRewriter 

11from renaissance.syntax_tree.match_finder import PatternMatch 

12from renaissance.syntax_tree.node_protocol import NodeProtocol 

13from renaissance.syntax_tree.semantic_kind import SemanticKind 

14 

15 

16class ASTProcessor: 

17 """AI: Processor that coordinates parsing, pattern matching, and rewriting of a single AST.""" 

18 

19 def __init__( 

20 self, 

21 root: ASTNode, 

22 ast_factory: ASTFactory, 

23 in_memory: bool = False, 

24 ) -> None: 

25 """AI: Prepare a processor that runs refactoring/analysis steps over the given AST.""" 

26 self.__root_node = root 

27 self.__rewriter = ASTRewriter(root) 

28 self.__ast_factory = ast_factory 

29 self.in_memory = in_memory 

30 self.repeat_step = 0 

31 

32 @property 

33 def factory(self) -> ASTFactory: 

34 """AI: Return the AST factory used to create nodes for this processor.""" 

35 return self.__ast_factory 

36 

37 @property 

38 def node(self) -> ASTNode: 

39 """AI: Return the root AST node this processor operates on.""" 

40 return self.__root_node 

41 

42 @property 

43 def filename(self) -> str: 

44 """AI: Return the filename of the AST this processor operates on.""" 

45 return self.__rewriter.get_filename() 

46 

47 @property 

48 def root(self) -> ASTNode: 

49 """AI: Return the root AST node this processor operates on.""" 

50 return self.__root_node 

51 

52 def replace( 

53 self, 

54 new_content: str, 

55 target: ASTNode | Sequence[ASTNode] | PatternMatch | Sequence[PatternMatch], 

56 include_whitespace: bool = True, 

57 include_comments: bool = True, 

58 ) -> None: 

59 """AI: Queue a rewrite replacing target's source text with new_content.""" 

60 self.__rewriter.replace(new_content, target, include_whitespace, include_comments) 

61 

62 def remove( 

63 self, 

64 target: ASTNode | Sequence[ASTNode] | PatternMatch | Sequence[PatternMatch], 

65 include_whitespace: bool = True, 

66 include_comments: bool = True, 

67 ) -> None: 

68 """AI: Queue a rewrite removing target's source text.""" 

69 self.__rewriter.remove(target, include_whitespace, include_comments) 

70 

71 def insert_before( 

72 self, 

73 new_content: str, 

74 target: ASTNode | Sequence[ASTNode] | PatternMatch | Sequence[PatternMatch], 

75 include_whitespace: bool = True, 

76 include_comments: bool = True, 

77 ) -> None: 

78 """AI: Queue a rewrite inserting new_content immediately before target's source text.""" 

79 self.__rewriter.insert_before(new_content, target, include_whitespace, include_comments) 

80 

81 def insert_after( 

82 self, 

83 new_content: str, 

84 target: ASTNode | Sequence[ASTNode] | PatternMatch | Sequence[PatternMatch], 

85 include_whitespace: bool = True, 

86 include_comments: bool = True, 

87 ) -> None: 

88 """AI: Queue a rewrite inserting new_content immediately after target's source text.""" 

89 self.__rewriter.insert_after(new_content, target, include_whitespace, include_comments) 

90 

91 def find_all(self, function: Callable[[ASTNode], Iterator[ASTNode] | bool]) -> Sequence[ASTNode]: 

92 """AI: Return all descendant nodes of the root for which function returns a truthy value or child iterator.""" 

93 return ASTFinder.find_all(self.__root_node, function) 

94 

95 def find_semantic_kind(self, kind: SemanticKind) -> Sequence[NodeProtocol]: 

96 """AI: Return all descendant nodes of the root matching the given semantic kind.""" 

97 return find_semantic_kind(self.__root_node, kind) 

98 

99 def find_match(self, *patterns_list, recursive: bool = True) -> Sequence[PatternMatch]: 

100 """AI: Return matches of patterns_list against the root's children.""" 

101 return renaissance.syntax_tree.match_finder.find_all( 

102 self.__root_node.children, 

103 *patterns_list, 

104 recursive=recursive, 

105 ) 

106 

107 def has_changed(self) -> bool: 

108 """AI: Return whether any rewrites have been queued.""" 

109 return self.__rewriter.has_changed() 

110 

111 def apply_to_string(self) -> str: 

112 """AI: Return the rewritten source as a string, applying all queued rewrites.""" 

113 return self.__rewriter.apply_to_string() 

114 

115 def commit(self) -> ASTProcessor: 

116 """Commit the current changes to the AST (Abstract Syntax Tree) and return a new ASTProcessor instance. 

117 

118 This method applies the current changes to the source code and creates a new ASTProcessor instance 

119 with the updated AST. If the changes are in-memory, it directly creates the new AST from the updated 

120 code string. Otherwise, it writes the changes to the file, reloads the file, and then creates the new AST. 

121 

122 Returns: 

123 ASTProcessor: A new instance of ASTProcessor with the updated AST. 

124 

125 Raises: 

126 IOError: If there is an error writing to the file. 

127 

128 """ 

129 if not self.__rewriter.has_changed(): 

130 return self 

131 self.__root_node, self.__rewriter = self._commit(self.__rewriter, self.__ast_factory, self.in_memory) 

132 return ASTProcessor(self.__root_node, self.__ast_factory, self.in_memory) 

133 

134 @staticmethod 

135 def _commit(rewriter: ASTRewriter, factory: ASTFactory, in_memory: bool = False): 

136 rewriter.apply_to_string() 

137 if in_memory: 

138 atu = factory.create_from_text(rewriter.apply_to_string(), rewriter.get_filename()) 

139 return atu, ASTRewriter(atu) 

140 # save file first then reload it 

141 with Path(rewriter.get_filename()).open("wb") as f: 

142 f.write(rewriter.apply()) 

143 atu = factory.create(Path(rewriter.get_filename())) 

144 return atu, ASTRewriter(atu) 

145 

146 

147# main 

148if __name__ == "__main__": 

149 

150 def test[T](_: str, factory: type[T]) -> T: 

151 """AI: Smoke-test helper verifying a factory callable returns an instance of its own type.""" 

152 result = factory() 

153 assert isinstance(result, factory) 

154 return result 

155 

156 test("key", str)