Files
code-memory/main.py
jze9 3f392bffca Fix streaming TransferEncodingError
DB saves were happening inside the generator after [DONE], causing
connection abort if any await failed. Moved to BackgroundTask.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-05-28 20:35:03 +05:00

215 lines
8.6 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"]
# ── models list endpoint ──────────────────────────────────────────────────────
@app.get("/v1/models")
async def list_models():
return {
"object": "list",
"data": [{
"id": "code-memory",
"object": "model",
"owned_by": "code-memory",
"created": 1700000000,
}]
}
# ── 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:
collected: list[str] = []
async def gen() -> AsyncIterator[bytes]:
try:
async for chunk in oll.chat_stream(messages, model=req.model):
collected.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"
except Exception as e:
yield f"data: {json.dumps({'error': str(e)})}\n\n".encode()
async def _save_stream():
response_text = "".join(collected)
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_stream)
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"}