tolivert's picture
deploy: financial_rag streamlit app
4e316d6
Raw
History Blame Contribute Delete
5.79 kB
"""
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,
)