niishantth/codegraph
0
1"""Multi-language AST parser using tree-sitter."""2 3from __future__ import annotations4 5import os6from pathlib import Path7from typing import Any, Optional8 9from tree_sitter import Language, Parser, Tree, Node10 11from codegraph.models import Entity, EntityType, Position, Span12 13# Language bindings14LANG_MAP: dict[str, Language] = {}15 16try:17 import tree_sitter_python as tspython18 LANG_MAP["py"] = Language(tspython.language())19except Exception:20 pass21 22try:23 import tree_sitter_javascript as tsjs24 LANG_MAP["js"] = Language(tsjs.language())25except Exception:26 pass27 28try:29 import tree_sitter_go as tsgo30 LANG_MAP["go"] = Language(tsgo.language())31except Exception:32 pass33 34try:35 import tree_sitter_java as tsjava36 LANG_MAP["java"] = Language(tsjava.language())37except Exception:38 pass39 40 41def _get_language(file_path: str) -> Optional[Language]:42 ext = Path(file_path).suffix.lstrip(".")43 return LANG_MAP.get(ext)44 45 46def _node_span(node: Node) -> Span:47 return Span(48 start=Position(line=node.start_point[0] + 1, column=node.start_point[1]),49 end=Position(line=node.end_point[0] + 1, column=node.end_point[1]),50 )51 52 53def _make_id(file_path: str, entity_type: EntityType, name: str, line: int) -> str:54 return f"{file_path}::{entity_type.value}::{name}::{line}"55 56 57class SourceParser:58 """Parse a single source file into entities."""59 60 def __init__(self, file_path: str, source: str) -> None:61 self.file_path = file_path62 self.source = source63 self.language = _get_language(file_path)64 self.tree: Optional[Tree] = None65 self.entities: list[Entity] = []66 67 def parse(self) -> list[Entity]:68 if self.language is None:69 return []70 parser = Parser(self.language)71 self.tree = parser.parse(bytes(self.source, "utf8"))72 root = self.tree.root_node73 self._walk(root)74 return self.entities75 76 def _walk(self, node: Node) -> None:77 handler = getattr(self, f"_handle_{self.language.name}_{node.type}", None)78 if handler:79 handler(node)80 for child in node.children:81 self._walk(child)82 83 # ---------- Python handlers ----------84 85 def _handle_python_function_definition(self, node: Node) -> None:86 name_node = node.child_by_field_name("name")87 if name_node is None:88 return89 name = name_node.text.decode("utf8")90 entity = Entity(91 id=_make_id(self.file_path, EntityType.FUNCTION, name, name_node.start_point[0] + 1),92 type=EntityType.FUNCTION,93 name=name,94 file_path=self.file_path,95 span=_node_span(node),96 meta={"decorators": self._extract_decorators(node)},97 )98 self.entities.append(entity)99 # Capture nested definitions (methods inside classes handled by class)100 self._capture_body_entities(node, parent_id=entity.id)101 102 def _handle_python_class_definition(self, node: Node) -> None:103 name_node = node.child_by_field_name("name")104 if name_node is None:105 return106 name = name_node.text.decode("utf8")107 entity = Entity(108 id=_make_id(self.file_path, EntityType.CLASS, name, name_node.start_point[0] + 1),109 type=EntityType.CLASS,110 name=name,111 file_path=self.file_path,112 span=_node_span(node),113 meta={"bases": self._extract_bases(node)},114 )115 self.entities.append(entity)116 self._capture_body_entities(node, parent_id=entity.id)117 118 def _capture_body_entities(self, node: Node, parent_id: str) -> None:119 body = node.child_by_field_name("body")120 if body is None:121 return122 for child in body.children:123 if child.type == "function_definition":124 name_node = child.child_by_field_name("name")125 if name_node is None:126 continue127 name = name_node.text.decode("utf8")128 e_type = EntityType.METHOD if "class" in parent_id else EntityType.FUNCTION129 self.entities.append(130 Entity(131 id=_make_id(self.file_path, e_type, name, name_node.start_point[0] + 1),132 type=e_type,133 name=name,134 file_path=self.file_path,135 span=_node_span(child),136 meta={"parent_id": parent_id, "decorators": self._extract_decorators(child)},137 )138 )139 elif child.type == "class_definition":140 name_node = child.child_by_field_name("name")141 if name_node is None:142 continue143 name = name_node.text.decode("utf8")144 self.entities.append(145 Entity(146 id=_make_id(self.file_path, EntityType.CLASS, name, name_node.start_point[0] + 1),147 type=EntityType.CLASS,148 name=name,149 file_path=self.file_path,150 span=_node_span(child),151 meta={"parent_id": parent_id, "bases": self._extract_bases(child)},152 )153 )154 155 def _extract_decorators(self, node: Node) -> list[str]:156 decs = []157 for child in node.children:158 if child.type == "decorator":159 text = child.text.decode("utf8").strip()160 decs.append(text)161 return decs162 163 def _extract_bases(self, node: Node) -> list[str]:164 bases = []165 bases_node = node.child_by_field_name("superclasses")166 if bases_node is None:167 # tree-sitter-python sometimes uses argument_list for bases168 for child in node.children:169 if child.type == "argument_list":170 bases_node = child171 break172 if bases_node:173 for child in bases_node.children:174 if child.type in ("identifier", "attribute"):175 bases.append(child.text.decode("utf8"))176 return bases177 178 # ---------- JavaScript / TypeScript handlers ----------179 180 def _handle_javascript_function_declaration(self, node: Node) -> None:181 name_node = node.child_by_field_name("name")182 if name_node is None:183 return184 name = name_node.text.decode("utf8")185 self.entities.append(186 Entity(187 id=_make_id(self.file_path, EntityType.FUNCTION, name, name_node.start_point[0] + 1),188 type=EntityType.FUNCTION,189 name=name,190 file_path=self.file_path,191 span=_node_span(node),192 )193 )194 195 def _handle_javascript_class_declaration(self, node: Node) -> None:196 name_node = node.child_by_field_name("name")197 if name_node is None:198 return199 name = name_node.text.decode("utf8")200 self.entities.append(201 Entity(202 id=_make_id(self.file_path, EntityType.CLASS, name, name_node.start_point[0] + 1),203 type=EntityType.CLASS,204 name=name,205 file_path=self.file_path,206 span=_node_span(node),207 )208 )209 210 def _handle_javascript_method_definition(self, node: Node) -> None:211 name_node = node.child_by_field_name("name")212 if name_node is None:213 return214 name = name_node.text.decode("utf8")215 self.entities.append(216 Entity(217 id=_make_id(self.file_path, EntityType.METHOD, name, name_node.start_point[0] + 1),218 type=EntityType.METHOD,219 name=name,220 file_path=self.file_path,221 span=_node_span(node),222 )223 )224 225 # ---------- Go handlers ----------226 227 def _handle_go_function_declaration(self, node: Node) -> None:228 name_node = node.child_by_field_name("name")229 if name_node is None:230 return231 name = name_node.text.decode("utf8")232 self.entities.append(233 Entity(234 id=_make_id(self.file_path, EntityType.FUNCTION, name, name_node.start_point[0] + 1),235 type=EntityType.FUNCTION,236 name=name,237 file_path=self.file_path,238 span=_node_span(node),239 )240 )241 242 def _handle_go_method_declaration(self, node: Node) -> None:243 name_node = node.child_by_field_name("name")244 if name_node is None:245 return246 name = name_node.text.decode("utf8")247 self.entities.append(248 Entity(249 id=_make_id(self.file_path, EntityType.METHOD, name, name_node.start_point[0] + 1),250 type=EntityType.METHOD,251 name=name,252 file_path=self.file_path,253 span=_node_span(node),254 )255 )256 257 # ---------- Java handlers ----------258 259 def _handle_java_class_declaration(self, node: Node) -> None:260 name_node = node.child_by_field_name("name")261 if name_node is None:262 return263 name = name_node.text.decode("utf8")264 self.entities.append(265 Entity(266 id=_make_id(self.file_path, EntityType.CLASS, name, name_node.start_point[0] + 1),267 type=EntityType.CLASS,268 name=name,269 file_path=self.file_path,270 span=_node_span(node),271 )272 )273 274 def _handle_java_method_declaration(self, node: Node) -> None:275 name_node = node.child_by_field_name("name")276 if name_node is None:277 return278 name = name_node.text.decode("utf8")279 self.entities.append(280 Entity(281 id=_make_id(self.file_path, EntityType.METHOD, name, name_node.start_point[0] + 1),282 type=EntityType.METHOD,283 name=name,284 file_path=self.file_path,285 span=_node_span(node),286 )287 )288 289 290class CodebaseScanner:291 """Scan a directory tree and parse every supported source file."""292 293 EXTENSIONS = {".py", ".js", ".jsx", ".ts", ".tsx", ".go", ".java"}294 295 def __init__(self, root: str, ignore_dirs: Optional[set[str]] = None) -> None:296 self.root = Path(root).resolve()297 self.ignore_dirs = ignore_dirs or {".git", "__pycache__", "node_modules", ".venv", "venv", ".tox", "build", "dist"}298 self.files: list[str] = []299 self.all_entities: list[Entity] = []300 301 def scan(self) -> list[Entity]:302 for path in self.root.rglob("*"):303 if path.is_dir() and path.name in self.ignore_dirs:304 continue305 if path.is_file() and path.suffix in self.EXTENSIONS:306 self.files.append(str(path))307 for fp in self.files:308 try:309 source = Path(fp).read_text(encoding="utf-8", errors="ignore")310 except Exception:311 continue312 parser = SourceParser(fp, source)313 entities = parser.parse()314 # Add a FILE entity for every parsed file315 file_entity = Entity(316 id=f"{fp}::file::{Path(fp).name}::1",317 type=EntityType.FILE,318 name=Path(fp).name,319 file_path=fp,320 span=Span(start=Position(line=1, column=0), end=Position(line=source.count("\n") + 1, column=0)),321 meta={"language": Path(fp).suffix.lstrip(".")},322 )323 self.all_entities.append(file_entity)324 self.all_entities.extend(entities)325 return self.all_entities326 