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
« 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."""
3from collections.abc import Callable, Iterator, Sequence
4from pathlib import Path
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
16class ASTProcessor:
17 """AI: Processor that coordinates parsing, pattern matching, and rewriting of a single AST."""
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
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
37 @property
38 def node(self) -> ASTNode:
39 """AI: Return the root AST node this processor operates on."""
40 return self.__root_node
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()
47 @property
48 def root(self) -> ASTNode:
49 """AI: Return the root AST node this processor operates on."""
50 return self.__root_node
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)
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)
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)
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)
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)
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)
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 )
107 def has_changed(self) -> bool:
108 """AI: Return whether any rewrites have been queued."""
109 return self.__rewriter.has_changed()
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()
115 def commit(self) -> ASTProcessor:
116 """Commit the current changes to the AST (Abstract Syntax Tree) and return a new ASTProcessor instance.
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.
122 Returns:
123 ASTProcessor: A new instance of ASTProcessor with the updated AST.
125 Raises:
126 IOError: If there is an error writing to the file.
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)
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)
147# main
148if __name__ == "__main__":
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
156 test("key", str)