OpenAI-compatible FastAPI proxy with GraphRAG + persistent memory. Includes 3-level Letta memory, CMA consolidation, AST-based code indexing. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
209 lines
8.9 KiB
Python
209 lines
8.9 KiB
Python
"""
|
|
GraphRAG: AST-based code graph stored in PostgreSQL.
|
|
Supports Python (ast module). Multi-hop traversal for architecture queries.
|
|
"""
|
|
import ast, os, json
|
|
from pathlib import Path
|
|
from uuid import UUID
|
|
|
|
from db import get_pool
|
|
from ollama_client import embed
|
|
|
|
def _vec_str(v: list[float]) -> str:
|
|
return "[" + ",".join(f"{x:.6f}" for x in v) + "]"
|
|
|
|
# ── AST Parser ────────────────────────────────────────────────────────────────
|
|
|
|
class PythonGraphBuilder(ast.NodeVisitor):
|
|
def __init__(self, file_path: str, source: str):
|
|
self.file_path = file_path
|
|
self.source_lines = source.splitlines()
|
|
self.nodes: list[dict] = []
|
|
self.edges: list[dict] = []
|
|
self._class_stack: list[str] = []
|
|
self._func_stack: list[str] = []
|
|
|
|
def _full_name(self, name: str) -> str:
|
|
parts = self._class_stack + self._func_stack + [name]
|
|
return ".".join(parts)
|
|
|
|
def _src(self, node) -> str:
|
|
try:
|
|
return ast.get_source_segment("\n".join(self.source_lines), node) or ""
|
|
except Exception:
|
|
return ""
|
|
|
|
def visit_ClassDef(self, node):
|
|
full = self._full_name(node.name)
|
|
self.nodes.append({"type": "class", "name": node.name, "full_name": full,
|
|
"file": self.file_path, "line_start": node.lineno,
|
|
"line_end": node.end_lineno, "content": self._src(node)})
|
|
self._class_stack.append(node.name)
|
|
self.generic_visit(node)
|
|
self._class_stack.pop()
|
|
|
|
def visit_FunctionDef(self, node):
|
|
full = self._full_name(node.name)
|
|
parent = ".".join(self._class_stack + self._func_stack) or None
|
|
self.nodes.append({"type": "function", "name": node.name, "full_name": full,
|
|
"file": self.file_path, "line_start": node.lineno,
|
|
"line_end": node.end_lineno, "content": self._src(node)})
|
|
if parent:
|
|
self.edges.append({"source": parent, "target": full, "rel": "has_method"})
|
|
self._func_stack.append(node.name)
|
|
self.generic_visit(node)
|
|
self._func_stack.pop()
|
|
|
|
visit_AsyncFunctionDef = visit_FunctionDef
|
|
|
|
def visit_Call(self, node):
|
|
caller = ".".join(self._class_stack + self._func_stack)
|
|
if not caller:
|
|
self.generic_visit(node)
|
|
return
|
|
if isinstance(node.func, ast.Attribute) and hasattr(node.func.value, "id"):
|
|
callee = f"{node.func.value.id}.{node.func.attr}"
|
|
elif isinstance(node.func, ast.Name):
|
|
callee = node.func.id
|
|
else:
|
|
self.generic_visit(node)
|
|
return
|
|
self.edges.append({"source": caller, "target": callee, "rel": "calls"})
|
|
self.generic_visit(node)
|
|
|
|
def visit_Import(self, node):
|
|
file_node = self.file_path
|
|
for alias in node.names:
|
|
self.edges.append({"source": file_node, "target": alias.name, "rel": "imports"})
|
|
self.generic_visit(node)
|
|
|
|
def visit_ImportFrom(self, node):
|
|
if node.module:
|
|
file_node = self.file_path
|
|
self.edges.append({"source": file_node, "target": node.module, "rel": "imports"})
|
|
self.generic_visit(node)
|
|
|
|
|
|
def parse_python_file(path: str) -> tuple[list[dict], list[dict]]:
|
|
try:
|
|
source = Path(path).read_text(encoding="utf-8", errors="ignore")
|
|
tree = ast.parse(source)
|
|
builder = PythonGraphBuilder(path, source)
|
|
builder.visit(tree)
|
|
# add file-level node
|
|
builder.nodes.insert(0, {"type": "file", "name": path, "full_name": path,
|
|
"file": path, "line_start": 1,
|
|
"line_end": len(source.splitlines()),
|
|
"content": source[:500]})
|
|
return builder.nodes, builder.edges
|
|
except SyntaxError:
|
|
return [], []
|
|
|
|
|
|
# ── Indexer ───────────────────────────────────────────────────────────────────
|
|
|
|
async def index_project(project_id: str, root_dir: str,
|
|
extensions: list[str] = None):
|
|
if extensions is None:
|
|
extensions = [".py"]
|
|
pool = await get_pool()
|
|
# clear old data
|
|
async with pool.acquire() as conn:
|
|
await conn.execute("DELETE FROM code_edges WHERE project_id=$1", project_id)
|
|
await conn.execute("DELETE FROM code_nodes WHERE project_id=$1", project_id)
|
|
|
|
node_name_to_id: dict[str, UUID] = {}
|
|
|
|
for ext in extensions:
|
|
for path in Path(root_dir).rglob(f"*{ext}"):
|
|
if ".git" in str(path) or "venv" in str(path) or "__pycache__" in str(path):
|
|
continue
|
|
nodes, edges = parse_python_file(str(path))
|
|
for n in nodes:
|
|
content_for_embed = f"{n['type']} {n['full_name']}\n{n['content'][:800]}"
|
|
vec = await embed(content_for_embed)
|
|
async with pool.acquire() as conn:
|
|
uid = await conn.fetchval("""
|
|
INSERT INTO code_nodes
|
|
(project_id, node_type, name, full_name, content,
|
|
embedding, file_path, line_start, line_end)
|
|
VALUES ($1,$2,$3,$4,$5,$6::vector,$7,$8,$9)
|
|
RETURNING id
|
|
""", project_id, n["type"], n["name"], n["full_name"],
|
|
n["content"][:2000], _vec_str(vec),
|
|
n["file"], n.get("line_start"), n.get("line_end"))
|
|
node_name_to_id[n["full_name"]] = uid
|
|
|
|
# store edges (resolve names to IDs later in a second pass)
|
|
for e in edges:
|
|
async with pool.acquire() as conn:
|
|
src_id = node_name_to_id.get(e["source"])
|
|
tgt_id = node_name_to_id.get(e["target"])
|
|
if src_id and tgt_id:
|
|
await conn.execute("""
|
|
INSERT INTO code_edges (project_id, source_id, target_id, relation_type)
|
|
VALUES ($1,$2,$3,$4) ON CONFLICT DO NOTHING
|
|
""", project_id, src_id, tgt_id, e["rel"])
|
|
|
|
return {"indexed_nodes": len(node_name_to_id)}
|
|
|
|
|
|
# ── Graph search ──────────────────────────────────────────────────────────────
|
|
|
|
async def search_code(query: str, project_id: str, limit: int = 5) -> list[dict]:
|
|
"""Vector search over code nodes."""
|
|
pool = await get_pool()
|
|
vec = await embed(query)
|
|
async with pool.acquire() as conn:
|
|
rows = await conn.fetch("""
|
|
SELECT id, node_type, full_name, content, file_path, line_start,
|
|
1-(embedding<=>$1::vector) AS score
|
|
FROM code_nodes WHERE project_id=$2
|
|
ORDER BY embedding<=>$1::vector LIMIT $3
|
|
""", _vec_str(vec), project_id, limit)
|
|
return [dict(r) for r in rows]
|
|
|
|
|
|
async def get_neighbors(node_id: UUID, project_id: str, hops: int = 2) -> list[dict]:
|
|
"""Multi-hop graph traversal starting from node_id."""
|
|
pool = await get_pool()
|
|
async with pool.acquire() as conn:
|
|
# recursive CTE for graph traversal
|
|
rows = await conn.fetch("""
|
|
WITH RECURSIVE graph AS (
|
|
SELECT id, node_type, full_name, content, file_path, 0 AS depth
|
|
FROM code_nodes WHERE id=$1
|
|
UNION
|
|
SELECT n.id, n.node_type, n.full_name, n.content, n.file_path,
|
|
g.depth+1
|
|
FROM code_nodes n
|
|
JOIN code_edges e ON (e.target_id=n.id AND e.source_id=g.id)
|
|
OR (e.source_id=n.id AND e.target_id=g.id)
|
|
JOIN graph g ON TRUE
|
|
WHERE g.depth < $2 AND n.project_id=$3
|
|
)
|
|
SELECT DISTINCT id, node_type, full_name, content, file_path, depth
|
|
FROM graph ORDER BY depth
|
|
LIMIT 30
|
|
""", node_id, hops, project_id)
|
|
return [dict(r) for r in rows]
|
|
|
|
|
|
async def graphrag_context(query: str, project_id: str) -> str:
|
|
"""Full GraphRAG: vector search + multi-hop expansion."""
|
|
seed_nodes = await search_code(query, project_id, limit=3)
|
|
if not seed_nodes:
|
|
return ""
|
|
all_nodes: dict[str, dict] = {str(n["id"]): n for n in seed_nodes}
|
|
for n in seed_nodes:
|
|
neighbors = await get_neighbors(n["id"], project_id, hops=2)
|
|
for nb in neighbors:
|
|
all_nodes[str(nb["id"])] = nb
|
|
parts = []
|
|
for n in sorted(all_nodes.values(), key=lambda x: x.get("depth", 0)):
|
|
parts.append(
|
|
f"[{n['node_type']}] {n['full_name']} ({n.get('file_path','')}"
|
|
f":{n.get('line_start','')})\n{(n.get('content') or '')[:400]}"
|
|
)
|
|
return "\n\n".join(parts[:10])
|