""" FastAPI application with chat, RAG, and document ingestion endpoints. Routes: - POST /chat — streaming chat via SSE (or JSON if stream=false) - POST /chat/rag — RAG-augmented chat (retrieve context, then generate) - POST /ingest — upload .txt/.md files, chunk + embed + store - GET /health — model status, device, chunk count The LLM runs synchronously in PyTorch, so generation is wrapped in run_in_executor to avoid blocking the async event loop. An asyncio Lock serialises generation calls since PyTorch is not thread-safe. """ import asyncio from contextlib import asynccontextmanager from fastapi import FastAPI, UploadFile, File from fastapi.responses import JSONResponse from sse_starlette.sse import EventSourceResponse from src.guardrails import GuardrailsMiddleware, PIICheck, ToxicityCheck, Severity from src.serving.engine import GenerationEngine from src.serving.models import ( ChatRequest, HealthResponse, IngestResponse, RAGChatRequest, ) from src.serving.rag import RAGPipeline # Global state populated at startup _engine: GenerationEngine | None = None _rag: RAGPipeline | None = None _lock: asyncio.Lock | None = None _guardrails: GuardrailsMiddleware | None = None @asynccontextmanager async def lifespan(app: FastAPI): global _engine, _rag, _lock, _guardrails model_size = app.state.model_size if hasattr(app.state, "model_size") else "0.6B" print(f"Loading GenerationEngine (Qwen3-{model_size})...") _engine = GenerationEngine(model_size=model_size) print("Loading RAG pipeline (BGE-small-en-v1.5)...") _rag = RAGPipeline() _lock = asyncio.Lock() print("Loading guardrails (PII + toxicity)...") _guardrails = GuardrailsMiddleware([ PIICheck(severity=Severity.WARN), ToxicityCheck(), ]) print("Ready.") yield print("Shutting down.") app = FastAPI(title="LLM Hub API", lifespan=lifespan) def _messages_to_dicts(messages) -> list[dict]: return [{"role": m.role, "content": m.content} for m in messages] async def _generate_sse(messages: list[dict], request: ChatRequest): """Wrap the sync generator into an async SSE stream.""" loop = asyncio.get_event_loop() async def event_generator(): async with _lock: gen = _engine.generate_stream( messages, max_tokens=request.max_tokens, temperature=request.temperature, top_k=request.top_k, top_p=request.top_p, ) while True: chunk = await loop.run_in_executor(None, lambda: next(gen, None)) if chunk is None: break yield {"data": chunk} return EventSourceResponse(event_generator()) async def _generate_full(messages: list[dict], request: ChatRequest): """Collect all tokens and return as a single JSON response.""" loop = asyncio.get_event_loop() def _run(): parts = [] gen = _engine.generate_stream( messages, max_tokens=request.max_tokens, temperature=request.temperature, top_k=request.top_k, top_p=request.top_p, ) for chunk in gen: parts.append(chunk) return "".join(parts) async with _lock: content = await loop.run_in_executor(None, _run) return JSONResponse({"role": "assistant", "content": content}) @app.post("/chat") async def chat(request: ChatRequest): messages = _messages_to_dicts(request.messages) # Scan the latest user message. if _guardrails and messages: last_user = next((m["content"] for m in reversed(messages) if m["role"] == "user"), "") scan = _guardrails.scan_input(last_user) if scan.blocked: return JSONResponse( status_code=400, content={"error": "Request blocked by safety filters.", "flags": [f.description for f in scan.flags]}, ) if request.stream: return await _generate_sse(messages, request) return await _generate_full(messages, request) @app.post("/chat/rag") async def chat_rag(request: RAGChatRequest): messages = _messages_to_dicts(request.messages) # Scan the latest user message. if _guardrails and messages: last_user = next((m["content"] for m in reversed(messages) if m["role"] == "user"), "") scan = _guardrails.scan_input(last_user) if scan.blocked: return JSONResponse( status_code=400, content={"error": "Request blocked by safety filters.", "flags": [f.description for f in scan.flags]}, ) augmented = _rag.build_rag_messages(messages, top_k=request.top_k_docs) if request.stream: return await _generate_sse(augmented, request) return await _generate_full(augmented, request) @app.post("/ingest") async def ingest(file: UploadFile = File(...)): if not file.filename.endswith((".txt", ".md")): return JSONResponse( status_code=400, content={"error": "Only .txt and .md files are supported"}, ) content = (await file.read()).decode("utf-8") num_chunks = _rag.ingest(content) return IngestResponse( filename=file.filename, num_chunks=num_chunks, message=f"Ingested {num_chunks} chunks from {file.filename}", ) @app.get("/health") async def health(): return HealthResponse( status="ok" if _engine is not None else "loading", model=f"Qwen3-{_engine.model_size}" if _engine else "none", device=str(_engine.device) if _engine else "none", num_chunks=_rag.store.num_chunks if _rag else 0, )