ai / backend /graph /dependency_graph.py
3v324v23's picture
agent
ee9a09c
Raw
History Blame Contribute Delete
5.87 kB
import ast
import json
import re
from pathlib import Path
from backend.core.constants import GRAPH_DIR
from backend.database.schemas import Chunk
IMPORT_RE = re.compile(r"^\s*import\s+(.+)|^\s*from\s+([.\w]+)\s+import\s+(.+)")
class LocalNetworkXGraphStore:
"""Build dependency graphs with NetworkX and persist them as local JSON files."""
def __init__(self) -> None:
try:
import networkx as nx
from networkx.readwrite import json_graph
except ImportError as exc:
raise ImportError("networkx is required for local graph storage.") from exc
self.nx = nx
self.json_graph = json_graph
self.graph_dir = GRAPH_DIR
self.graph_dir.mkdir(parents=True, exist_ok=True)
self._graphs: dict[str, self.nx.DiGraph] = {}
def build(self, repo_id: str, repo_path: Path, chunks: list[Chunk]) -> None:
graph = self.nx.DiGraph(repo_id=repo_id, repo_path=str(repo_path))
files = sorted({chunk.path for chunk in chunks})
for file_path in files:
graph.add_node(
self._file_node(file_path),
kind="file",
path=file_path,
label=file_path,
)
for chunk in chunks:
file_node = self._file_node(chunk.path)
if chunk.kind != "imports":
chunk_node = self._chunk_node(chunk.id)
graph.add_node(
chunk_node,
kind=chunk.kind,
chunk_id=chunk.id,
path=chunk.path,
language=chunk.language,
symbol=chunk.symbol,
start_line=chunk.start_line,
end_line=chunk.end_line,
label=chunk.symbol,
)
graph.add_edge(file_node, chunk_node, relation="contains")
for imported in self._extract_imports(chunk):
import_node = self._import_node(imported)
graph.add_node(
import_node,
kind="import",
name=imported,
label=imported,
)
graph.add_edge(file_node, import_node, relation="imports")
self._graphs[repo_id] = graph
self._save(repo_id, graph)
def neighbors_for_terms(self, repo_id: str, terms: list[str]) -> set[str]:
if not terms:
return set()
graph = self._graphs.get(repo_id) or self._load(repo_id)
if not graph:
return set()
lowered_terms = [term.lower() for term in terms]
matches = {
node
for node, attrs in graph.nodes(data=True)
if self._node_matches(attrs, lowered_terms)
}
related: set[str] = set()
for node in matches:
self._add_search_values(related, graph.nodes[node])
for neighbor in set(graph.predecessors(node)) | set(graph.successors(node)):
self._add_search_values(related, graph.nodes[neighbor])
return related
def _save(self, repo_id: str, graph) -> None:
payload = self.json_graph.node_link_data(graph)
self._graph_path(repo_id).write_text(json.dumps(payload, indent=2), encoding="utf-8")
def _load(self, repo_id: str):
graph_path = self._graph_path(repo_id)
if not graph_path.exists():
return None
payload = json.loads(graph_path.read_text(encoding="utf-8"))
graph = self.json_graph.node_link_graph(payload, directed=True)
self._graphs[repo_id] = graph
return graph
def _graph_path(self, repo_id: str) -> Path:
safe_repo_id = re.sub(r"[^A-Za-z0-9_.-]+", "_", repo_id)
return self.graph_dir / f"{safe_repo_id}.json"
def _node_matches(self, attrs: dict, lowered_terms: list[str]) -> bool:
searchable_values = [
attrs.get("path", ""),
attrs.get("symbol", ""),
attrs.get("kind", ""),
attrs.get("language", ""),
attrs.get("name", ""),
attrs.get("label", ""),
]
haystack = " ".join(str(value).lower() for value in searchable_values if value)
return any(term in haystack for term in lowered_terms)
def _add_search_values(self, related: set[str], attrs: dict) -> None:
for key in ("path", "symbol", "name", "label"):
value = attrs.get(key)
if value:
related.add(str(value))
def _extract_imports(self, chunk: Chunk) -> list[str]:
if chunk.language == "python":
try:
tree = ast.parse(chunk.content)
except SyntaxError:
return []
imports: list[str] = []
for node in ast.walk(tree):
if isinstance(node, ast.Import):
imports.extend(alias.name for alias in node.names)
elif isinstance(node, ast.ImportFrom) and node.module:
imports.append(node.module)
return imports
imports = []
for line in chunk.content.splitlines():
match = IMPORT_RE.match(line)
if match:
imports.append(next(group for group in match.groups() if group))
return imports
def _file_node(self, path: str) -> str:
return f"file:{path}"
def _chunk_node(self, chunk_id: str) -> str:
return f"chunk:{chunk_id}"
def _import_node(self, name: str) -> str:
return f"import:{name}"
# Backup: Neo4jGraphStore was previously selected when USE_NEO4J=true.
# The active implementation now always uses local NetworkX JSON graph storage,
# so .env can keep legacy Neo4j values without affecting runtime.
dependency_graph_store = LocalNetworkXGraphStore()