Coverage for src/renaissance/syntax_tree/ast_finder.py: 69%

45 statements  

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

1"""AI: Helpers for finding AST nodes matching a predicate or semantic kind.""" 

2 

3import re 

4from collections.abc import Callable, Iterator, Sequence 

5 

6from renaissance.utils.ast_utils import traverse 

7 

8from .ast_node import ASTNode 

9from .node_protocol import NodeProtocol 

10from .semantic_kind import SemanticKind 

11 

12 

13class ASTFinder: 

14 """AI: Static helpers for finding AST nodes matching a predicate or semantic kind.""" 

15 

16 KIND_MATCH = re.compile(r"[\W_]+") 

17 

18 @staticmethod 

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

20 """AI: Return all descendant nodes of ast_node for which function returns a truthy value or child iterator.""" 

21 return list(ASTFinder.__find_all(ast_node, function)) 

22 

23 # @staticmethod 

24 # def find_kind(ast_node: ASTNode, kind: str | re.Pattern[str]) -> Sequence[ASTNode]: 

25 # return list(ASTFinder.__matches_kind(ast_node, kind)) 

26 

27 @staticmethod 

28 def find(ast_node: ASTNode, kind: str | re.Pattern[str]) -> Sequence[ASTNode]: 

29 """AI: Return all descendant nodes of ast_node whose kind matches the given kind pattern.""" 

30 return list(ASTFinder.__matches_kind(ast_node, kind)) 

31 

32 @staticmethod 

33 def matches_kind(ast_node: ASTNode | None, kind: str | re.Pattern[str]) -> bool: 

34 """AI: Return whether ast_node's kind matches the given kind pattern.""" 

35 # compare kind with the ast_node kind only using word characters 

36 # get kind of the ast_node with only word characters 

37 if ast_node is None: 

38 return False 

39 ast_kind = ASTFinder.KIND_MATCH.sub("", ast_node.kind).lower() 

40 pattern = kind if isinstance(kind, re.Pattern) else re.compile(kind, re.IGNORECASE) 

41 return pattern.fullmatch(ast_kind) is not None 

42 

43 @staticmethod 

44 def __find_all(ast_node: ASTNode, function: Callable[[ASTNode], Iterator[ASTNode] | bool]) -> Iterator[ASTNode]: 

45 result = function(ast_node) 

46 if isinstance(result, bool) and result: 

47 yield ast_node 

48 elif isinstance(result, Iterator): 

49 yield from result 

50 for child in ast_node.children: 

51 yield from ASTFinder.__find_all(child, function) 

52 

53 @staticmethod 

54 def __matches_kind(ast_node: ASTNode, kind: str | re.Pattern[str]) -> Iterator[ASTNode]: 

55 pattern = kind if isinstance(kind, re.Pattern) else re.compile(kind, re.IGNORECASE) 

56 node_kind = ast_node.kind or "" 

57 ast_kind = ASTFinder.KIND_MATCH.sub("", node_kind).lower() 

58 

59 if pattern.fullmatch(ast_kind): 

60 yield ast_node 

61 for child in ast_node.children: 

62 # assert isinstance(child, type(ast_node)), f'Expected {type(ast_node)} but got {type(child)}' 

63 yield from ASTFinder.__matches_kind(child, pattern) 

64 

65 

66def find_nodes(ast_node: NodeProtocol, predicate) -> Sequence[NodeProtocol]: 

67 """AI: Return all descendants (and ast_node itself) matching predicate, via a full traversal.""" 

68 return [node for node in traverse(ast_node) if predicate(node)] 

69 

70 

71def matches_node(ast_node: NodeProtocol, predicate) -> bool: 

72 """AI: Return True if ast_node itself satisfies predicate.""" 

73 return predicate(ast_node) 

74 

75 

76def find_semantic_kind(ast_node: NodeProtocol, kind: SemanticKind) -> Sequence[NodeProtocol]: 

77 """AI: Return all nodes under ast_node whose semantic kind matches the given kind.""" 

78 return find_nodes(ast_node, lambda node: node.semantic_kind is kind)