Coverage for src/renaissance/syntax_tree/ast_refactor_actions.py: 88%
60 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: Higher-level refactoring actions built on top of pattern matching and rewriting."""
3from collections.abc import Callable, Sequence
4from functools import cache
5from typing import TYPE_CHECKING
7from .ast_finder import matches_node
8from .ast_node import ASTNode
9from .ast_processor import ASTProcessor
10from .match_finder import MatchFinder, PatternMatch
11from .semantic_kind import SemanticKind
13if TYPE_CHECKING:
14 from renaissance.integrations.clang.c_pattern_factory import CPPPatternFactory
17def _kind_predicate(kind):
18 if kind is None:
19 return lambda _node: False
20 if callable(kind):
21 return kind
22 if isinstance(kind, SemanticKind):
23 return lambda node: node.semantic_kind is kind
24 raise TypeError("kind must be a SemanticKind or a node predicate")
27class ASTRefactorActions:
28 """AI: Higher-level refactoring actions (replace, insert, remove) built on top of pattern matching and rewriting."""
30 def __init__(self, processor: ASTProcessor, pattern_factory: CPPPatternFactory) -> None:
31 """AI: Provide pattern-based refactoring actions (replace, insert, remove) over an AST."""
32 self.processor = processor
33 self.pattern_factory = pattern_factory
34 self.replaced: set[int] = set()
36 def replace_expr(self, name: str, replacement: str, kind: SemanticKind | Callable[[ASTNode], bool]):
37 """AI: Replace occurrences of the named expression matching kind with replacement."""
38 kind_predicate = _kind_predicate(kind)
40 def test(n: ASTNode):
41 if (kind and matches_node(n, kind_predicate)) and n.name == name:
42 yield n
44 [self.processor.replace(found.text.replace(found.name, replacement, 1), found) for found in self.processor.find_all(test)]
46 def replace_name(
47 self,
48 name: str,
49 replacement: str,
50 kind: SemanticKind | Callable[[ASTNode], bool] | None = None,
51 skip_kind: SemanticKind | Callable[[ASTNode], bool] | None = None,
52 ):
53 """AI: Replace occurrences of the named node (matching kind, excluding skip_kind) with replacement."""
54 kind_predicate = _kind_predicate(kind)
55 skip_kind_predicate = _kind_predicate(skip_kind)
57 def matches_name(n1: ASTNode) -> bool:
58 return (
59 (kind is None or matches_node(n1, kind_predicate))
60 and (skip_kind is None or not matches_node(n1, skip_kind_predicate))
61 and n1
62 and n1.name == name
63 )
65 found_nodes = self.processor.find_all(matches_name)
66 [self.replaced.add(found.offset) for found in found_nodes if found.offset not in self.replaced]
67 for n in found_nodes:
68 self.processor.replace(n.text.replace(n.name, replacement, 1), n)
70 def replace_text(
71 self,
72 text: str,
73 replacement: str,
74 kind: SemanticKind | Callable[[ASTNode], bool] | None = None,
75 skip_kind: SemanticKind | Callable[[ASTNode], bool] | None = None,
76 ):
77 """AI: Replace nodes whose text equals text (matching kind, excluding skip_kind) with replacement."""
78 kind_predicate = _kind_predicate(kind)
79 skip_kind_predicate = _kind_predicate(skip_kind)
81 def matches_text(n: ASTNode) -> bool:
82 return (
83 (kind is None or matches_node(n, kind_predicate))
84 and (skip_kind is None or not matches_node(n, skip_kind_predicate))
85 and n is not None
86 and n.text == text
87 )
89 found_nodes = self.processor.find_all(matches_text)
90 [self.replaced.add(found.offset) for found in found_nodes if found.offset not in self.replaced]
92 [self.processor.replace(n.text.replace(n.name, replacement, 1), n) for n in found_nodes]
94 def replace_declaration(self, declaration: str, replacement: str):
95 """AI: Replace every match of declaration with replacement."""
96 for match in self.find_declaration(declaration):
97 self.processor.replace(replacement, match)
99 def _replace_patterns(
100 self,
101 node: ASTNode,
102 replacement: str,
103 patterns: Sequence[Sequence[ASTNode]],
104 matches: Sequence[PatternMatch],
105 ):
106 if not patterns:
107 self.processor.replace(replacement, matches)
108 return
109 [
110 self._replace_patterns(m.nodes[0], replacement, patterns[1:], list(matches) + [m])
111 for m in MatchFinder.find_all([node], patterns[0])
112 ]
114 @cache
115 def find_declaration(self, decl_pattern: str):
116 """AI: Return the matches of decl_pattern parsed as a declaration pattern."""
117 pattern = self.pattern_factory.create_declaration(decl_pattern)
118 return self.processor.find_match(pattern)
120 @cache
121 def collect(self, pattern: str, pattern_kind: str):
122 """AI: Return the matches of pattern parsed as the given pattern_kind."""
123 root = self.pattern_factory.create(pattern, pattern_kind)
125 return self.processor.find_match(root)