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>
195 lines
8.0 KiB
Python
195 lines
8.0 KiB
Python
"""
|
|
code-memory: OpenAI-compatible proxy with GraphRAG + persistent memory.
|
|
Endpoints:
|
|
POST /v1/chat/completions — drop-in Ollama replacement with memory
|
|
POST /v1/index — index a code repository
|
|
GET /v1/memories — list memories
|
|
DELETE /v1/memories/{id} — delete a memory
|
|
GET /health — health check
|
|
"""
|
|
import json, os, asyncio
|
|
from contextlib import asynccontextmanager
|
|
from typing import AsyncIterator
|
|
|
|
import asyncpg
|
|
from fastapi import FastAPI, HTTPException, BackgroundTasks
|
|
from fastapi.responses import StreamingResponse
|
|
from pydantic import BaseModel, Field
|
|
from dotenv import load_dotenv
|
|
|
|
load_dotenv()
|
|
|
|
import db as _db
|
|
import memory as mem
|
|
import graph as gph
|
|
import ollama_client as oll
|
|
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
|
|
|
# ── lifespan ──────────────────────────────────────────────────────────────────
|
|
|
|
scheduler = AsyncIOScheduler()
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(app: FastAPI):
|
|
await _db.init_db()
|
|
# CMA consolidation runs every 6 hours
|
|
scheduler.add_job(mem.consolidate_memories, "interval", hours=6,
|
|
kwargs={"user_id": "default"})
|
|
scheduler.start()
|
|
yield
|
|
scheduler.shutdown()
|
|
|
|
app = FastAPI(title="code-memory", lifespan=lifespan)
|
|
|
|
# ── models ────────────────────────────────────────────────────────────────────
|
|
|
|
class Message(BaseModel):
|
|
role: str
|
|
content: str
|
|
|
|
class ChatRequest(BaseModel):
|
|
model: str = Field(default=None)
|
|
messages: list[Message]
|
|
stream: bool = False
|
|
session_id: str = "default"
|
|
user_id: str = "default"
|
|
project_id: str | None = None
|
|
|
|
class IndexRequest(BaseModel):
|
|
project_id: str
|
|
root_dir: str
|
|
extensions: list[str] = [".py"]
|
|
|
|
# ── context assembly ──────────────────────────────────────────────────────────
|
|
|
|
async def build_context(request: ChatRequest) -> list[dict]:
|
|
user_msg = next((m.content for m in reversed(request.messages)
|
|
if m.role == "user"), "")
|
|
|
|
# 1. Relevant archival memories
|
|
memories = await mem.search_memories(user_msg, user_id=request.user_id, limit=6)
|
|
mem_block = ""
|
|
if memories:
|
|
mem_block = "## Relevant memories\n" + "\n".join(
|
|
f"- [{m['memory_type']}] {m['content']}" for m in memories
|
|
if m.get("score", 0) > 0.4
|
|
)
|
|
|
|
# 2. GraphRAG code context
|
|
code_block = ""
|
|
if request.project_id:
|
|
ctx = await gph.graphrag_context(user_msg, request.project_id)
|
|
if ctx:
|
|
code_block = f"## Code architecture context\n{ctx}"
|
|
|
|
# 3. Conversation summary (if exists)
|
|
summary = await mem.get_summary(request.session_id)
|
|
summary_block = f"## Earlier conversation summary\n{summary}" if summary else ""
|
|
|
|
# 4. Assemble system prompt
|
|
base_system = (
|
|
"You are an expert coding assistant with deep knowledge of the project architecture. "
|
|
"Use the provided memory and code context to give precise, informed answers. "
|
|
"Always reason about the full architecture, not just isolated snippets."
|
|
)
|
|
context_parts = [p for p in [base_system, summary_block, mem_block, code_block] if p]
|
|
system_content = "\n\n".join(context_parts)
|
|
|
|
# 5. Build messages: system + recent history + current
|
|
recent = await mem.get_recent_turns(request.session_id, limit=20)
|
|
messages = [{"role": "system", "content": system_content}]
|
|
messages += [{"role": r["role"], "content": r["content"]} for r in recent]
|
|
# add current user messages (skip if already in history)
|
|
for m in request.messages:
|
|
if not any(h["role"] == m.role and h["content"] == m.content for h in recent[-2:]):
|
|
messages.append({"role": m.role, "content": m.content})
|
|
|
|
return messages
|
|
|
|
# ── chat endpoint ─────────────────────────────────────────────────────────────
|
|
|
|
@app.post("/v1/chat/completions")
|
|
async def chat_completions(req: ChatRequest, bg: BackgroundTasks):
|
|
messages = await build_context(req)
|
|
user_content = next((m.content for m in reversed(req.messages)
|
|
if m.role == "user"), "")
|
|
|
|
if req.stream:
|
|
async def gen() -> AsyncIterator[bytes]:
|
|
full_response = []
|
|
async for chunk in oll.chat_stream(messages, model=req.model):
|
|
full_response.append(chunk)
|
|
data = json.dumps({"choices": [{"delta": {"content": chunk},
|
|
"finish_reason": None}]})
|
|
yield f"data: {data}\n\n".encode()
|
|
yield b"data: [DONE]\n\n"
|
|
# save after stream
|
|
response_text = "".join(full_response)
|
|
await mem.save_turn(req.session_id, "user", user_content)
|
|
await mem.save_turn(req.session_id, "assistant", response_text)
|
|
await mem.extract_and_store(user_content, user_id=req.user_id)
|
|
await mem.extract_and_store(response_text, user_id=req.user_id)
|
|
await mem.summarize_old_turns(req.session_id)
|
|
|
|
return StreamingResponse(gen(), media_type="text/event-stream")
|
|
|
|
response_text = await oll.chat(messages, model=req.model)
|
|
|
|
# background: save turns + extract facts
|
|
async def _save():
|
|
await mem.save_turn(req.session_id, "user", user_content)
|
|
await mem.save_turn(req.session_id, "assistant", response_text)
|
|
await mem.extract_and_store(user_content, user_id=req.user_id)
|
|
await mem.extract_and_store(response_text, user_id=req.user_id)
|
|
await mem.summarize_old_turns(req.session_id)
|
|
|
|
bg.add_task(_save)
|
|
|
|
return {
|
|
"id": "cmpl-1",
|
|
"object": "chat.completion",
|
|
"model": req.model or os.getenv("CHAT_MODEL"),
|
|
"choices": [{"index": 0,
|
|
"message": {"role": "assistant", "content": response_text},
|
|
"finish_reason": "stop"}]
|
|
}
|
|
|
|
# ── index endpoint ────────────────────────────────────────────────────────────
|
|
|
|
@app.post("/v1/index")
|
|
async def index_code(req: IndexRequest, bg: BackgroundTasks):
|
|
if not os.path.isdir(req.root_dir):
|
|
raise HTTPException(404, f"Directory not found: {req.root_dir}")
|
|
bg.add_task(gph.index_project, req.project_id, req.root_dir, req.extensions)
|
|
return {"status": "indexing started", "project_id": req.project_id}
|
|
|
|
# ── memory endpoints ──────────────────────────────────────────────────────────
|
|
|
|
@app.get("/v1/memories")
|
|
async def list_memories(user_id: str = "default", limit: int = 50):
|
|
pool = await _db.get_pool()
|
|
async with pool.acquire() as conn:
|
|
rows = await conn.fetch("""
|
|
SELECT id, content, importance, memory_type, access_count,
|
|
last_accessed, created_at
|
|
FROM memories WHERE user_id=$1
|
|
AND (metadata->>'archived') IS NULL
|
|
ORDER BY importance DESC, last_accessed DESC LIMIT $2
|
|
""", user_id, limit)
|
|
return [dict(r) for r in rows]
|
|
|
|
@app.delete("/v1/memories/{memory_id}")
|
|
async def delete_memory(memory_id: str):
|
|
pool = await _db.get_pool()
|
|
async with pool.acquire() as conn:
|
|
await conn.execute("DELETE FROM memories WHERE id=$1",
|
|
__import__("uuid").UUID(memory_id))
|
|
return {"deleted": memory_id}
|
|
|
|
@app.get("/health")
|
|
async def health():
|
|
pool = await _db.get_pool()
|
|
async with pool.acquire() as conn:
|
|
await conn.fetchval("SELECT 1")
|
|
return {"status": "ok", "db": "connected"}
|