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

1"""AI: Higher-level refactoring actions built on top of pattern matching and rewriting.""" 

2 

3from collections.abc import Callable, Sequence 

4from functools import cache 

5from typing import TYPE_CHECKING 

6 

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 

12 

13if TYPE_CHECKING: 

14 from renaissance.integrations.clang.c_pattern_factory import CPPPatternFactory 

15 

16 

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

25 

26 

27class ASTRefactorActions: 

28 """AI: Higher-level refactoring actions (replace, insert, remove) built on top of pattern matching and rewriting.""" 

29 

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

35 

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) 

39 

40 def test(n: ASTNode): 

41 if (kind and matches_node(n, kind_predicate)) and n.name == name: 

42 yield n 

43 

44 [self.processor.replace(found.text.replace(found.name, replacement, 1), found) for found in self.processor.find_all(test)] 

45 

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) 

56 

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 ) 

64 

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) 

69 

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) 

80 

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 ) 

88 

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] 

91 

92 [self.processor.replace(n.text.replace(n.name, replacement, 1), n) for n in found_nodes] 

93 

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) 

98 

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 ] 

113 

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) 

119 

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) 

124 

125 return self.processor.find_match(root)