Team Ai
Modelpublic

niishantth/codegraph

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
graph_builder.py214 linesDownload Raw Back to codegraph
1"""Build the knowledge graph from parsed entities and extracted relationships."""2 3from __future__ import annotations4 5import re6from pathlib import Path7from typing import Any, Optional8 9import networkx as nx10 11from codegraph.models import Entity, EntityType, RelationType, Relationship12from codegraph.parser import SourceParser13 14 15class GraphBuilder:16    """Construct a directed NetworkX graph from entities and relationships."""17 18    def __init__(self) -> None:19        self.graph: nx.DiGraph = nx.DiGraph()20        self.entity_map: dict[str, Entity] = {}21 22    def add_entities(self, entities: list[Entity]) -> None:23        for e in entities:24            self.entity_map[e.id] = e25            self.graph.add_node(e.id, **e.model_dump())26 27    def add_relationships(self, relationships: list[Relationship]) -> None:28        for r in relationships:29            if r.source in self.entity_map and r.target in self.entity_map:30                self.graph.add_edge(r.source, r.target, type=r.type.value, **r.meta)31 32    def build_from_entities(self, entities: list[Entity]) -> nx.DiGraph:33        """Infer relationships from entities and source text."""34        self.add_entities(entities)35        rels = RelationshipExtractor(entities).extract()36        self.add_relationships(rels)37        return self.graph38 39    def get_subgraph(self, node_id: str, depth: int = 2) -> nx.DiGraph:40        """Extract a neighbourhood subgraph around a node."""41        nodes = {node_id}42        current = {node_id}43        for _ in range(depth):44            next_layer: set[str] = set()45            for n in current:46                next_layer.update(self.graph.predecessors(n))47                next_layer.update(self.graph.successors(n))48            nodes.update(next_layer)49            current = next_layer50        return self.graph.subgraph(nodes).copy()51 52    def query_by_name(self, name: str) -> list[Entity]:53        results = []54        for e in self.entity_map.values():55            if e.name == name:56                results.append(e)57        return results58 59    def query_by_file(self, file_path: str) -> list[Entity]:60        return [e for e in self.entity_map.values() if e.file_path == file_path]61 62    def query_by_line(self, file_path: str, line: int) -> Optional[Entity]:63        """Return the innermost entity that contains the given line."""64        candidates: list[Entity] = []65        for e in self.entity_map.values():66            if e.file_path == file_path and e.span and e.span.start.line <= line <= e.span.end.line:67                candidates.append(e)68        if not candidates:69            return None70        # Prefer the smallest span (most specific)71        candidates.sort(key=lambda e: (e.span.end.line - e.span.start.line if e.span else 10**9))72        return candidates[0]73 74 75class RelationshipExtractor:76    """Extract relationships by re-parsing source text for imports, calls, inheritance."""77 78    def __init__(self, entities: list[Entity]) -> None:79        self.entities = entities80        self.entity_by_name: dict[str, list[Entity]] = {}81        for e in entities:82            self.entity_by_name.setdefault(e.name, []).append(e)83        self.file_entities: dict[str, list[Entity]] = {}84        for e in entities:85            self.file_entities.setdefault(e.file_path, []).append(e)86 87    def extract(self) -> list[Relationship]:88        rels: list[Relationship] = []89        for e in self.entities:90            if e.type == EntityType.FILE:91                continue92            # CONTAINS: file -> entity93            rels.append(94                Relationship(95                    source=e.file_path + f"::file::{Path(e.file_path).name}::1",96                    target=e.id,97                    type=RelationType.CONTAINS,98                )99            )100        # Language-specific extraction101        for file_path, ents in self.file_entities.items():102            try:103                source = Path(file_path).read_text(encoding="utf-8", errors="ignore")104            except Exception:105                continue106            rels.extend(self._extract_imports(file_path, source))107            rels.extend(self._extract_calls(file_path, source))108            rels.extend(self._extract_inheritance(file_path, source))109        return rels110 111    def _extract_imports(self, file_path: str, source: str) -> list[Relationship]:112        rels: list[Relationship] = []113        ext = Path(file_path).suffix114        file_entity_id = file_path + f"::file::{Path(file_path).name}::1"115        if ext == ".py":116            for match in re.finditer(r"^\s*(?:from\s+(\S+)\s+import|import\s+(\S+))", source, re.MULTILINE):117                mod = match.group(1) or match.group(2)118                if mod:119                    rels.append(120                        Relationship(121                            source=file_entity_id,122                            target=mod,123                            type=RelationType.IMPORTS,124                            meta={"raw": match.group(0).strip(), "line": source[:match.start()].count("\n") + 1},125                        )126                    )127        elif ext in (".js", ".jsx", ".ts", ".tsx"):128            for match in re.finditer(r"import\s+.*?\s+from\s+['\"]([^'\"]+)['\"]", source):129                rels.append(130                    Relationship(131                        source=file_entity_id,132                        target=match.group(1),133                        type=RelationType.IMPORTS,134                        meta={"raw": match.group(0), "line": source[:match.start()].count("\n") + 1},135                    )136                )137        elif ext == ".go":138            for match in re.finditer(r'import\s+\(?\s*(?:"([^"]+)"|\S+\s+"([^"]+)")', source):139                mod = match.group(1) or match.group(2)140                if mod:141                    rels.append(142                        Relationship(143                            source=file_entity_id,144                            target=mod,145                            type=RelationType.IMPORTS,146                            meta={"raw": match.group(0), "line": source[:match.start()].count("\n") + 1},147                        )148                    )149        elif ext == ".java":150            for match in re.finditer(r"import\s+([^;]+);", source):151                rels.append(152                    Relationship(153                        source=file_entity_id,154                        target=match.group(1).strip(),155                        type=RelationType.IMPORTS,156                        meta={"raw": match.group(0), "line": source[:match.start()].count("\n") + 1},157                    )158                )159        return rels160 161    def _extract_calls(self, file_path: str, source: str) -> list[Relationship]:162        rels: list[Relationship] = []163        # Simple regex-based call extraction for cross-file linking164        # Match identifier followed by '(' – crude but fast165        for match in re.finditer(r"(?<![\.\w])([a-zA-Z_]\w*)\s*\(", source):166            name = match.group(1)167            line = source[:match.start()].count("\n") + 1168            # Link to any entity with same name in the same file first, then globally169            targets = self.entity_by_name.get(name, [])170            same_file = [t for t in targets if t.file_path == file_path]171            chosen = same_file[0] if same_file else (targets[0] if targets else None)172            if chosen:173                # Find the function/method that contains this call174                container = self._find_container(file_path, line)175                if container:176                    rels.append(177                        Relationship(178                            source=container.id,179                            target=chosen.id,180                            type=RelationType.CALLS,181                            meta={"line": line, "call": name},182                        )183                    )184        return rels185 186    def _extract_inheritance(self, file_path: str, source: str) -> list[Relationship]:187        rels: list[Relationship] = []188        for e in self.file_entities.get(file_path, []):189            if e.type == EntityType.CLASS and e.meta.get("bases"):190                for base in e.meta["bases"]:191                    targets = self.entity_by_name.get(base, [])192                    same_file = [t for t in targets if t.file_path == file_path]193                    chosen = same_file[0] if same_file else (targets[0] if targets else None)194                    if chosen:195                        rels.append(196                            Relationship(197                                source=e.id,198                                target=chosen.id,199                                type=RelationType.INHERITS,200                            )201                        )202        return rels203 204    def _find_container(self, file_path: str, line: int) -> Optional[Entity]:205        candidates: list[Entity] = []206        for e in self.file_entities.get(file_path, []):207            if e.type in (EntityType.FUNCTION, EntityType.METHOD, EntityType.CLASS) and e.span:208                if e.span.start.line <= line <= e.span.end.line:209                    candidates.append(e)210        if not candidates:211            return None212        candidates.sort(key=lambda e: (e.span.end.line - e.span.start.line) if e.span else 10**9)213        return candidates[0]214