Coverage for src / renaissance / integrations / tree_sitter / extractor.py: 100%

86 statements  

« prev     ^ index     » next       coverage.py v7.13.4, created at 2026-09-09 14:04 +0000

1from pathlib import Path 

2 

3import networkx 

4 

5from renaissance.integrations.tree_sitter.adapter import TreeSitterAdapter 

6from renaissance.integrations.tree_sitter.factory import TreeSitterPatternFactory 

7from renaissance.integrations.types import Call, FunctionDef 

8from renaissance.syntax_tree import PatternMatch 

9from renaissance.syntax_tree.match_finder import match_pattern 

10 

11GRAPHML_DIR = "out_graphml" 

12Path(GRAPHML_DIR).mkdir(parents=True, exist_ok=True) 

13 

14 

15class Extractor: 

16 def __init__(self, factory: TreeSitterPatternFactory, patterns: list[str]): 

17 self.pattern_factory = factory 

18 self.patterns = patterns 

19 

20 def run(self, raw: str) -> list[PatternMatch]: 

21 code = self.pattern_factory.create_statements(raw) 

22 results = [] 

23 for rule in self.patterns: 

24 pattern = self.pattern_factory.create_statements(rule) 

25 results.extend(match_pattern(code, pattern, {})) 

26 return results 

27 

28 

29class BaseCodeGraphExtractor: 

30 def __init__(self, language: str, lib_path: str): 

31 self.language = language 

32 self.lib_path = lib_path 

33 self.adapter = TreeSitterAdapter(lib_path) 

34 self.graph = networkx.DiGraph() 

35 

36 def extract(self, files): 

37 for f in files: 

38 try: 

39 code = Path(f).read_text() 

40 tree = self.adapter.parse_code(code) 

41 lst = self.adapter.to_lst(code, tree) 

42 self._process_file(f, lst) 

43 except Exception as e: 

44 print(f"Error processing {f}: {e}") 

45 

46 def _process_file(self, file_path: str, lst): 

47 raise NotImplementedError 

48 

49 def save_graph(self, filename: str): 

50 path = Path(GRAPHML_DIR) / filename 

51 networkx.write_graphml(self.graph, path) 

52 print(f"Graph saved to: {path}") 

53 

54 

55class PythonCodeGraphExtractor(BaseCodeGraphExtractor): 

56 def _process_file(self, file_path, lst): 

57 folder = str(Path(file_path).parent) 

58 self.graph.add_node(file_path, type="file", folder=folder) 

59 self.graph.add_node(folder, type="folder") 

60 self.graph.add_edge(folder, file_path, type="contains") 

61 

62 for node in lst.traverse(): 

63 if node.ast_type == FunctionDef: 

64 name = node.signature.split("(")[0].split()[-1] 

65 self.graph.add_node(name, type="function", file=file_path) 

66 self.graph.add_edge(file_path, name, type="defines") 

67 

68 elif node.ast_type == Call: 

69 call_target = node.signature.strip().split("(")[0] 

70 self.graph.add_node(call_target, type="call_target") 

71 self.graph.add_edge(file_path, call_target, type="calls") 

72 

73 

74class JavaCodeGraphExtractor(BaseCodeGraphExtractor): 

75 def _process_file(self, file_path, lst): 

76 folder = str(Path(file_path).parent) 

77 self.graph.add_node(file_path, type="file", folder=folder) 

78 self.graph.add_node(folder, type="folder") 

79 self.graph.add_edge(folder, file_path, type="contains") 

80 

81 for node in lst.traverse(): 

82 if node.ast_type == FunctionDef: 

83 name = node.properties.get("name", "method") 

84 self.graph.add_node(name, type="method", file=file_path) 

85 self.graph.add_edge(file_path, name, type="defines") 

86 

87 elif node.ast_type == Call: 

88 target = node.signature.strip().split("(")[0] 

89 self.graph.add_node(target, type="method_target") 

90 self.graph.add_edge(file_path, target, type="calls") 

91 

92 

93class CppCodeGraphExtractor(BaseCodeGraphExtractor): 

94 def _process_file(self, file_path, lst): 

95 folder = str(Path(file_path).parent) 

96 self.graph.add_node(file_path, type="file", folder=folder) 

97 self.graph.add_node(folder, type="folder") 

98 self.graph.add_edge(folder, file_path, type="contains") 

99 

100 for node in lst.traverse(): 

101 if node.ast_type == FunctionDef: 

102 name = node.properties.get("name", "func") 

103 self.graph.add_node(name, type="function", file=file_path) 

104 self.graph.add_edge(file_path, name, type="defines") 

105 

106 elif node.ast_type == Call: 

107 call_expr = node.signature.strip().split("(")[0] 

108 self.graph.add_node(call_expr, type="call_target") 

109 self.graph.add_edge(file_path, call_expr, type="calls")