CoolFace
Apppublic

salim0986/graph-bug-ai

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
temporary_graph.py454 linesDownload Raw Back to src
1"""2Temporary GraphRAG Context Builder3Builds in-memory graph and vector context for PR files without permanent storage4"""5 6from typing import Dict, List, Set, Optional, Tuple, Any7from dataclasses import dataclass, field8from collections import defaultdict9from .logger import setup_logger10from .parser import UniversalParser11import hashlib12 13logger = setup_logger(__name__)14 15 16# ============================================================================17# DATA MODELS18# ============================================================================19 20@dataclass21class TempNode:22    """Temporary graph node (function, class, etc.)"""23    id: str24    name: str25    type: str  # function, class, method, etc.26    file: str27    line: int28    end_line: int29    code: str30    language: str31    32    # Relationships33    calls: Set[str] = field(default_factory=set)  # Node IDs this calls34    called_by: Set[str] = field(default_factory=set)  # Node IDs that call this35    imports: Set[str] = field(default_factory=set)  # Module/file imports36    37    def to_dict(self) -> Dict[str, Any]:38        """Convert to dict for context"""39        return {40            "id": self.id,41            "name": self.name,42            "type": self.type,43            "file": self.file,44            "line": self.line,45            "end_line": self.end_line,46            "code": self.code[:500],  # Truncate for context47            "language": self.language48        }49 50 51@dataclass52class TempFileNode:53    """Temporary file-level node"""54    path: str55    language: str56    nodes: List[TempNode] = field(default_factory=list)57    imports: Set[str] = field(default_factory=set)58    depends_on: Set[str] = field(default_factory=set)  # File dependencies59 60 61@dataclass62class TempVectorEntry:63    """Temporary vector entry for similarity search"""64    id: str65    text: str66    embedding: Optional[List[float]] = None67    metadata: Dict[str, Any] = field(default_factory=dict)68    69    def similarity(self, other: "TempVectorEntry") -> float:70        """Calculate cosine similarity with another entry"""71        if self.embedding is None or other.embedding is None:72            return 0.073        74        # Cosine similarity75        dot_product = sum(a * b for a, b in zip(self.embedding, other.embedding))76        magnitude_a = sum(a * a for a in self.embedding) ** 0.577        magnitude_b = sum(b * b for b in other.embedding) ** 0.578        79        if magnitude_a == 0 or magnitude_b == 0:80            return 0.081        82        return dot_product / (magnitude_a * magnitude_b)83 84 85# ============================================================================86# TEMPORARY GRAPH BUILDER87# ============================================================================88 89class TemporaryGraphBuilder:90    """91    Builds in-memory graph structure from PR files92    93    Features:94    - Parse files with tree-sitter95    - Extract nodes (functions, classes, methods)96    - Build call graphs and dependencies97    - Find relationships with existing code98    """99    100    def __init__(self, parser: UniversalParser):101        self.parser = parser102        self.nodes: Dict[str, TempNode] = {}  # node_id -> node103        self.files: Dict[str, TempFileNode] = {}  # file_path -> file_node104        self.node_index: Dict[str, List[str]] = defaultdict(list)  # name -> node_ids105        106    def process_file(self, filename: str, content: str, language: str) -> TempFileNode:107        """108        Parse and process a single file109        110        Args:111            filename: Relative file path112            content: File content113            language: Programming language114            115        Returns:116            TempFileNode with extracted nodes and relationships117        """118        logger.info(f"[TempGraph] Processing file: {filename} ({language})")119        120        # Create file node121        file_node = TempFileNode(path=filename, language=language)122        123        # Skip invalid/unsupported languages124        if not language or language == "text":125            logger.warning(f"[TempGraph] Invalid language '{language}' for {filename}, skipping")126            return file_node127        128        try:129            # Parse with tree-sitter130            content_bytes = content.encode('utf-8')131            captures, _ = self.parser.parse_code(content_bytes, language)132            133            if not captures:134                logger.warning(f"[TempGraph] No captures found for {filename}")135                return file_node136            137            # Extract nodes from captures138            for node, capture_name in captures:139                temp_node = self._create_temp_node(node, capture_name, filename, language, content_bytes)140                if temp_node:141                    self.nodes[temp_node.id] = temp_node142                    file_node.nodes.append(temp_node)143                    self.node_index[temp_node.name].append(temp_node.id)144            145            # Extract imports146            file_node.imports = self._extract_imports(content, language)147            148            # Build intra-file call graph149            self._build_call_graph(file_node)150            151            self.files[filename] = file_node152            logger.info(f"[TempGraph] Processed {filename}: {len(file_node.nodes)} nodes, {len(file_node.imports)} imports")153            154        except Exception as e:155            logger.error(f"[TempGraph] Error processing {filename}: {e}", exc_info=True)156        157        return file_node158    159    def _create_temp_node(160        self,161        tree_node: Any,162        capture_name: str,163        filename: str,164        language: str,165        content_bytes: bytes166    ) -> Optional[TempNode]:167        """Create a TempNode from tree-sitter node"""168        try:169            # Filter by node type (only functions, classes, methods)170            node_type = tree_node.type171            if node_type not in [172                "function_definition", "function_declaration", "function_item",173                "method_definition", "method_declaration",174                "class_definition", "class_declaration"175            ]:176                return None177            178            # Extract code179            code = content_bytes[tree_node.start_byte:tree_node.end_byte].decode('utf-8', errors='ignore')180            181            # Extract name (first line, cleaned)182            first_line = code.split('\n')[0].strip()183            name = first_line[:100]184            185            # Generate unique ID186            node_id = hashlib.md5(f"{filename}:{tree_node.start_point[0]}:{name}".encode()).hexdigest()[:16]187            188            return TempNode(189                id=node_id,190                name=name,191                type=node_type,192                file=filename,193                line=tree_node.start_point[0],194                end_line=tree_node.end_point[0],195                code=code,196                language=language197            )198            199        except Exception as e:200            logger.error(f"[TempGraph] Error creating temp node: {e}")201            return None202    203    def _extract_imports(self, content: str, language: str) -> Set[str]:204        """Extract import statements from file content"""205        imports = set()206        207        try:208            lines = content.split('\n')209            210            # Language-specific import patterns211            if language in ['python', 'py']:212                for line in lines:213                    line = line.strip()214                    if line.startswith('import ') or line.startswith('from '):215                        # Extract module name216                        if line.startswith('import '):217                            module = line.replace('import ', '').split()[0].strip()218                        else:219                            module = line.split('from ')[1].split('import')[0].strip()220                        imports.add(module)221            222            elif language in ['javascript', 'typescript', 'jsx', 'tsx']:223                for line in lines:224                    line = line.strip()225                    if 'import' in line and 'from' in line:226                        # Extract module path227                        parts = line.split('from')228                        if len(parts) > 1:229                            module = parts[1].strip().strip(';').strip('"').strip("'")230                            imports.add(module)231            232        except Exception as e:233            logger.error(f"[TempGraph] Error extracting imports: {e}")234        235        return imports236    237    def _build_call_graph(self, file_node: TempFileNode):238        """Build call relationships within a file"""239        # Simple heuristic: look for function calls in code240        for node in file_node.nodes:241            for other_node in file_node.nodes:242                if node.id != other_node.id:243                    # Check if node's code mentions other_node's name244                    if other_node.name in node.code:245                        node.calls.add(other_node.id)246                        other_node.called_by.add(node.id)247    248    def build_dependencies(self):249        """Build file-level dependency graph"""250        for file_path, file_node in self.files.items():251            for import_path in file_node.imports:252                # Check if import corresponds to another file in the temporary graph253                if import_path in self.files:254                    file_node.depends_on.add(import_path)255    256    def get_node(self, node_id: str) -> Optional[TempNode]:257        """Get node by ID"""258        return self.nodes.get(node_id)259    260    def get_nodes_by_name(self, name: str) -> List[TempNode]:261        """Get all nodes with a given name"""262        node_ids = self.node_index.get(name, [])263        return [self.nodes[nid] for nid in node_ids if nid in self.nodes]264    265    def get_file_dependencies(self, filename: str) -> List[str]:266        """Get files that the given file depends on"""267        file_node = self.files.get(filename)268        if not file_node:269            return []270        return list(file_node.depends_on)271    272    def get_node_dependencies(self, node_id: str) -> List[Dict[str, Any]]:273        """Get nodes that the given node calls"""274        node = self.nodes.get(node_id)275        if not node:276            return []277        278        return [279            self.nodes[called_id].to_dict()280            for called_id in node.calls281            if called_id in self.nodes282        ]283    284    def to_context_dict(self) -> Dict[str, Any]:285        """Convert temporary graph to context dictionary"""286        return {287            "files": {path: {288                "nodes": [n.to_dict() for n in fn.nodes],289                "imports": list(fn.imports),290                "dependencies": list(fn.depends_on)291            } for path, fn in self.files.items()},292            "total_nodes": len(self.nodes),293            "total_files": len(self.files)294        }295 296 297# ============================================================================298# TEMPORARY VECTOR BUILDER299# ============================================================================300 301class TemporaryVectorBuilder:302    """303    Builds in-memory vector index for PR files304    305    Features:306    - Generate embeddings using sentence-transformers307    - Similarity search308    - Find similar code across temporary and permanent contexts309    """310    311    def __init__(self, embed_model):312        self.embed_model = embed_model313        self.vectors: List[TempVectorEntry] = []314        self.index: Dict[str, int] = {}  # node_id -> vector index315        316    def add_nodes(self, nodes: List[TempNode]):317        """318        Add multiple nodes with batch embedding (OPTIMIZED)319        This is 10-20x faster than adding nodes one-by-one320        """321        if not nodes:322            return323        324        try:325            # Prepare texts for batch encoding326            texts = []327            node_list = []328            329            for node in nodes:330                # Create text for embedding (name + code snippet)331                text = f"{node.name}\n{node.code[:500]}"332                texts.append(text)333                node_list.append(node)334            335            # Batch encode all texts at once (MAJOR PERFORMANCE BOOST)336            logger.debug(f"[TempVector] Batch encoding {len(texts)} nodes...")337            embeddings = self.embed_model.encode(338                texts,339                batch_size=50,  # Process 50 at once340                show_progress_bar=False,341                convert_to_numpy=True342            )343            344            # Create vector entries345            for node, text, embedding in zip(node_list, texts, embeddings):346                entry = TempVectorEntry(347                    id=node.id,348                    text=text,349                    embedding=embedding.tolist(),350                    metadata={351                        "name": node.name,352                        "type": node.type,353                        "file": node.file,354                        "line": node.line,355                        "language": node.language356                    }357                )358                359                self.index[node.id] = len(self.vectors)360                self.vectors.append(entry)361            362            logger.debug(f"[TempVector] Batch encoded {len(embeddings)} nodes successfully")363            364        except Exception as e:365            logger.error(f"[TempVector] Error in batch encoding: {e}")366    367    def add_node(self, node: TempNode):368        """Add a single node (fallback for individual additions)"""369        try:370            # Create text for embedding (name + code snippet)371            text = f"{node.name}\n{node.code[:500]}"372            373            # Generate embedding374            embedding = self.embed_model.encode(text).tolist()375            376            # Create vector entry377            entry = TempVectorEntry(378                id=node.id,379                text=text,380                embedding=embedding,381                metadata={382                    "name": node.name,383                    "type": node.type,384                    "file": node.file,385                    "line": node.line,386                    "language": node.language387                }388            )389            390            self.index[node.id] = len(self.vectors)391            self.vectors.append(entry)392            393        except Exception as e:394            logger.error(f"[TempVector] Error adding node {node.id}: {e}")395    396    def search_similar(self, query: str, limit: int = 5, min_score: float = 0.7) -> List[Dict[str, Any]]:397        """Search for similar code using text query"""398        try:399            # Generate query embedding400            query_embedding = self.embed_model.encode(query).tolist()401            402            # Create temporary entry for similarity calculation403            query_entry = TempVectorEntry(id="query", text=query, embedding=query_embedding)404            405            # Calculate similarities406            results = []407            for entry in self.vectors:408                score = query_entry.similarity(entry)409                if score >= min_score:410                    results.append({411                        "score": score,412                        "node_id": entry.id,413                        "metadata": entry.metadata,414                        "text": entry.text[:200]415                    })416            417            # Sort by score and limit418            results.sort(key=lambda x: x["score"], reverse=True)419            return results[:limit]420            421        except Exception as e:422            logger.error(f"[TempVector] Error searching: {e}")423            return []424    425    def find_similar_to_node(self, node_id: str, limit: int = 5, min_score: float = 0.7) -> List[Dict[str, Any]]:426        """Find nodes similar to a given node"""427        if node_id not in self.index:428            return []429        430        idx = self.index[node_id]431        source_entry = self.vectors[idx]432        433        results = []434        for entry in self.vectors:435            if entry.id != node_id:436                score = source_entry.similarity(entry)437                if score >= min_score:438                    results.append({439                        "score": score,440                        "node_id": entry.id,441                        "metadata": entry.metadata,442                        "text": entry.text[:200]443                    })444        445        results.sort(key=lambda x: x["score"], reverse=True)446        return results[:limit]447    448    def to_context_dict(self) -> Dict[str, Any]:449        """Convert temporary vectors to context dictionary"""450        return {451            "total_vectors": len(self.vectors),452            "indexed_nodes": list(self.index.keys())453        }454