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
« 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."""
3import re
4from collections.abc import Callable, Iterator, Sequence
6from renaissance.utils.ast_utils import traverse
8from .ast_node import ASTNode
9from .node_protocol import NodeProtocol
10from .semantic_kind import SemanticKind
13class ASTFinder:
14 """AI: Static helpers for finding AST nodes matching a predicate or semantic kind."""
16 KIND_MATCH = re.compile(r"[\W_]+")
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))
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))
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))
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
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)
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()
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)
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)]
71def matches_node(ast_node: NodeProtocol, predicate) -> bool:
72 """AI: Return True if ast_node itself satisfies predicate."""
73 return predicate(ast_node)
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)