""" 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"}