Initial commit: code-memory service
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>
This commit is contained in:
194
main.py
Normal file
194
main.py
Normal file
@@ -0,0 +1,194 @@
|
||||
"""
|
||||
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"}
|
||||
Reference in New Issue
Block a user