Team Ai
Modelpublic

niishantth/codegraph

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
parser.py326 linesDownload Raw Back to codegraph
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