| """Tree-sitter code symbols, structural retrieval, and reference-graph expansion.""" |
|
|
| from __future__ import annotations |
|
|
| from collections import Counter, defaultdict |
| from dataclasses import dataclass |
| from difflib import SequenceMatcher |
| import math |
| from typing import Iterable, Sequence |
|
|
| from tree_sitter import Language, Node, Parser |
| import tree_sitter_go |
| import tree_sitter_python |
|
|
| from .components import Candidate |
| from .repository import GitSnapshot |
| from .retrieval import query_terms, tokenize |
|
|
|
|
| GO_LANGUAGE = Language(tree_sitter_go.language()) |
| PYTHON_LANGUAGE = Language(tree_sitter_python.language()) |
| GO_DECLARATIONS = { |
| "function_declaration", |
| "method_declaration", |
| "type_spec", |
| "const_spec", |
| "var_spec", |
| } |
| PYTHON_DECLARATIONS = {"function_definition", "class_definition"} |
| GO_IDENTIFIERS = {"identifier", "field_identifier", "type_identifier", "package_identifier"} |
| PYTHON_IDENTIFIERS = {"identifier"} |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class CodeSymbol: |
| key: str |
| path: str |
| name: str |
| kind: str |
| line_start: int |
| line_end: int |
| signature: str |
| text: str |
| identifiers: tuple[str, ...] |
|
|
|
|
| def walk(node: Node) -> Iterable[Node]: |
| stack = [node] |
| while stack: |
| current = stack.pop() |
| yield current |
| stack.extend(reversed(current.children)) |
|
|
|
|
| def node_text(node: Node, source: bytes) -> str: |
| return source[node.start_byte : node.end_byte].decode("utf-8", errors="replace") |
|
|
|
|
| def _parse_file( |
| path: str, |
| text: str, |
| language: Language, |
| declarations: set[str], |
| identifier_types: set[str], |
| signature_delimiter: str, |
| ) -> tuple[CodeSymbol, ...]: |
| source = text.encode("utf-8") |
| parser = Parser(language) |
| tree = parser.parse(source) |
| symbols: list[CodeSymbol] = [] |
| occurrences: Counter[str] = Counter() |
| for node in walk(tree.root_node): |
| if node.type not in declarations: |
| continue |
| name_node = node.child_by_field_name("name") |
| if name_node is None: |
| continue |
| name = node_text(name_node, source) |
| body = node_text(node, source) |
| signature = ( |
| body.split(signature_delimiter, 1)[0].strip().splitlines()[0] |
| if body.strip() |
| else name |
| ) |
| occurrence = occurrences[name] |
| occurrences[name] += 1 |
| suffix = "" if occurrence == 0 else f"#{occurrence + 1}" |
| key = f"{path}::{name}{suffix}" |
| referenced_identifiers = tuple( |
| dict.fromkeys( |
| node_text(descendant, source) |
| for descendant in walk(node) |
| if descendant.type in identifier_types |
| ) |
| ) |
| symbols.append( |
| CodeSymbol( |
| key=key, |
| path=path, |
| name=name, |
| kind=node.type, |
| line_start=node.start_point[0] + 1, |
| line_end=node.end_point[0] + 1, |
| signature=signature, |
| text=body, |
| identifiers=referenced_identifiers, |
| ) |
| ) |
| return tuple(symbols) |
|
|
|
|
| def parse_go_file(path: str, text: str) -> tuple[CodeSymbol, ...]: |
| return _parse_file( |
| path, |
| text, |
| GO_LANGUAGE, |
| GO_DECLARATIONS, |
| GO_IDENTIFIERS, |
| "{", |
| ) |
|
|
|
|
| def parse_python_file(path: str, text: str) -> tuple[CodeSymbol, ...]: |
| return _parse_file( |
| path, |
| text, |
| PYTHON_LANGUAGE, |
| PYTHON_DECLARATIONS, |
| PYTHON_IDENTIFIERS, |
| ":", |
| ) |
|
|
|
|
| def parse_source_file(path: str, text: str, language: str) -> tuple[CodeSymbol, ...]: |
| if language == "go": |
| return parse_go_file(path, text) |
| if language == "python": |
| return parse_python_file(path, text) |
| raise ValueError(f"unsupported Tree-sitter language: {language}") |
|
|
|
|
| def parse_snapshot( |
| snapshot: GitSnapshot, |
| commit: str, |
| language: str = "go", |
| ) -> tuple[CodeSymbol, ...]: |
| suffixes = {"go": (".go",), "python": (".py",)} |
| if language not in suffixes: |
| raise ValueError(f"unsupported Tree-sitter language: {language}") |
| symbols: list[CodeSymbol] = [] |
| for source in snapshot.iter_files(commit, suffixes[language]): |
| symbols.extend(parse_source_file(source.path, source.text, language)) |
| return tuple(symbols) |
|
|
|
|
| |
| GoSymbol = CodeSymbol |
|
|
|
|
| class SyntaxRetriever: |
| """BM25 over declaration-level Tree-sitter representations.""" |
|
|
| def __init__(self, symbols: Sequence[CodeSymbol], k1: float = 1.2, b: float = 0.75): |
| self.symbols = tuple(symbols) |
| self.k1 = k1 |
| self.b = b |
| self.term_frequencies = tuple( |
| Counter(tokenize(f"{item.path}\n{item.kind} {item.name}\n{item.signature}\n{item.text}")) |
| for item in symbols |
| ) |
| self.lengths = tuple(sum(value.values()) for value in self.term_frequencies) |
| self.average_length = sum(self.lengths) / max(len(self.lengths), 1) |
| document_frequency: Counter[str] = Counter() |
| for frequencies in self.term_frequencies: |
| document_frequency.update(frequencies.keys()) |
| self.document_frequency = document_frequency |
|
|
| def retrieve(self, query: str, limit: int) -> Sequence[Candidate]: |
| terms = query_terms(query) |
| count = len(self.symbols) |
| ranked: list[tuple[float, CodeSymbol]] = [] |
| for symbol, frequencies, length in zip(self.symbols, self.term_frequencies, self.lengths): |
| score = 0.0 |
| for term in terms: |
| frequency = frequencies.get(term, 0) |
| if not frequency: |
| continue |
| df = self.document_frequency[term] |
| inverse_frequency = math.log(1.0 + (count - df + 0.5) / (df + 0.5)) |
| denominator = frequency + self.k1 * ( |
| 1.0 - self.b + self.b * length / max(self.average_length, 1.0) |
| ) |
| score += inverse_frequency * frequency * (self.k1 + 1.0) / denominator |
| fuzzy = max( |
| (SequenceMatcher(None, term, token).ratio() for term in terms for token in tokenize(symbol.name)), |
| default=0.0, |
| ) |
| if fuzzy >= 0.72: |
| score += (fuzzy - 0.72) * 4.0 |
| if score > 0.0: |
| ranked.append((score, symbol)) |
| ranked.sort(key=lambda item: (-item[0], item[1].path, item[1].line_start, item[1].name)) |
| return tuple( |
| Candidate( |
| path=symbol.path, |
| line_start=symbol.line_start, |
| line_end=symbol.line_end, |
| text=symbol.text, |
| source="tree_sitter_symbol", |
| score=score, |
| symbol=symbol.key, |
| metadata={"kind": symbol.kind, "signature": symbol.signature}, |
| ) |
| for score, symbol in ranked[:limit] |
| ) |
|
|
|
|
| class SymbolGraph: |
| """Undirected static reference graph between declaration symbols.""" |
|
|
| def __init__(self, symbols: Sequence[CodeSymbol]): |
| self.by_key = {item.key: item for item in symbols} |
| self.by_path: dict[str, list[CodeSymbol]] = defaultdict(list) |
| by_name: dict[str, list[CodeSymbol]] = defaultdict(list) |
| for symbol in symbols: |
| self.by_path[symbol.path].append(symbol) |
| by_name[symbol.name].append(symbol) |
| adjacency: dict[str, set[str]] = {item.key: set() for item in symbols} |
| for symbol in symbols: |
| for identifier in symbol.identifiers: |
| targets = by_name.get(identifier, ()) |
| |
| |
| |
| if not targets or len(targets) > 8: |
| continue |
| for target in targets: |
| if target.key == symbol.key: |
| continue |
| adjacency[symbol.key].add(target.key) |
| adjacency[target.key].add(symbol.key) |
| self.adjacency = {key: tuple(sorted(values)) for key, values in adjacency.items()} |
|
|
| def seed_keys(self, candidate: Candidate) -> tuple[str, ...]: |
| if candidate.symbol in self.by_key: |
| return (candidate.symbol,) |
| overlaps = [ |
| symbol.key |
| for symbol in self.by_path.get(candidate.path, ()) |
| if symbol.line_start <= candidate.line_end and candidate.line_start <= symbol.line_end |
| ] |
| if overlaps: |
| return tuple(overlaps) |
| return tuple(symbol.key for symbol in self.by_path.get(candidate.path, ())[:3]) |
|
|
| def expand( |
| self, |
| candidates: Sequence[Candidate], |
| hops: int, |
| limit: int, |
| seeds: int = 20, |
| neighbors_per_seed: int = 10, |
| ) -> tuple[Candidate, ...]: |
| if hops not in {1, 2}: |
| raise ValueError("graph expansion hops must be one or two") |
| scored: dict[str, tuple[float, Candidate]] = {} |
| frontier: dict[str, float] = {} |
| for rank, candidate in enumerate(candidates, start=1): |
| score = 1.0 / rank |
| file_key = candidate.path |
| if file_key not in scored or score > scored[file_key][0]: |
| scored[file_key] = (score, candidate) |
| if rank <= seeds: |
| for key in self.seed_keys(candidate): |
| frontier[key] = max(frontier.get(key, 0.0), score) |
|
|
| visited = set(frontier) |
| for hop in range(1, hops + 1): |
| next_frontier: dict[str, float] = {} |
| for parent_key, parent_score in sorted(frontier.items()): |
| for neighbor_key in self.adjacency.get(parent_key, ())[:neighbors_per_seed]: |
| if neighbor_key in visited: |
| continue |
| score = parent_score * (0.85**hop) |
| next_frontier[neighbor_key] = max(next_frontier.get(neighbor_key, 0.0), score) |
| symbol = self.by_key[neighbor_key] |
| candidate = Candidate( |
| path=symbol.path, |
| line_start=symbol.line_start, |
| line_end=symbol.line_end, |
| text=symbol.text, |
| source=f"graph_hop_{hop}", |
| score=score, |
| symbol=symbol.key, |
| metadata={"kind": symbol.kind, "signature": symbol.signature}, |
| ) |
| if symbol.path not in scored or score > scored[symbol.path][0]: |
| scored[symbol.path] = (score, candidate) |
| visited.update(next_frontier) |
| frontier = next_frontier |
| ranked = sorted( |
| scored.values(), |
| key=lambda item: (-item[0], item[1].path, item[1].line_start), |
| ) |
| return tuple(candidate for _, candidate in ranked[:limit]) |
|
|
|
|
| def symbol_hits( |
| candidates: Sequence[Candidate], |
| gold_symbols: Sequence[str], |
| symbols: Sequence[CodeSymbol], |
| cutoff: int = 10, |
| ) -> set[str]: |
| by_gold = {gold.split("#", 1)[0]: gold for gold in gold_symbols} |
| hits: set[str] = set() |
| symbol_lookup = {symbol.key.split("#", 1)[0]: symbol for symbol in symbols} |
| for candidate in candidates[:cutoff]: |
| candidate_key = (candidate.symbol or "").split("#", 1)[0] |
| if candidate_key in by_gold: |
| hits.add(by_gold[candidate_key]) |
| for normalized, gold in by_gold.items(): |
| symbol = symbol_lookup.get(normalized) |
| if ( |
| symbol is not None |
| and symbol.path == candidate.path |
| and symbol.line_start <= candidate.line_end |
| and candidate.line_start <= symbol.line_end |
| ): |
| hits.add(gold) |
| return hits |
|
|