Spaces:
Sleeping
Sleeping
| """ | |
| 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 | |
| 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}) | |
| 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) | |
| 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) | |
| 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}", | |
| ) | |
| 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, | |
| ) | |