import json from contextlib import asynccontextmanager from fastapi import FastAPI, HTTPException from fastapi.staticfiles import StaticFiles from fastapi.responses import HTMLResponse, StreamingResponse from pydantic import BaseModel from src.rag import answer, retrieve from src.reranker import _get_model @asynccontextmanager async def lifespan(app: FastAPI): print("Preloading reranker model...") _get_model() print("Reranker ready.") yield app = FastAPI(title="Vietnamese RAG API", lifespan=lifespan) class QueryRequest(BaseModel): question: str top_k: int = 5 class Source(BaseModel): source: str page: int rrf_score: float rerank_score: float has_context: bool = False text: str = "" class QueryResponse(BaseModel): answer: str sources: list[Source] rewritten_query: str = "" @app.get("/", response_class=HTMLResponse) def index(): with open("static/index.html", encoding="utf-8") as f: return f.read() @app.post("/query", response_model=QueryResponse) def query(req: QueryRequest): result = answer(req.question, top_k=req.top_k) return QueryResponse( answer=result["answer"], sources=[Source(**s) for s in result["sources"]], rewritten_query=result.get("rewritten_query", ""), ) @app.post("/stream") async def stream(req: QueryRequest): import asyncio from src.rag import build_prompt, client, LLM_MODEL async def generate(): # Padding 2KB để force nginx flush buffer yield ": " + " " * 2048 + "\n\n" yield f"data: {json.dumps({'type': 'thinking'})}\n\n" await asyncio.sleep(0) loop = asyncio.get_event_loop() contexts, sources, rewritten = await loop.run_in_executor(None, retrieve, req.question, req.top_k) yield f"data: {json.dumps({'type': 'sources', 'sources': sources, 'rewritten_query': rewritten})}\n\n" await asyncio.sleep(0) prompt = build_prompt(req.question, contexts) stream_resp = await loop.run_in_executor( None, lambda: client.chat.completions.create( model=LLM_MODEL, messages=[{"role": "user", "content": prompt}], stream=True, ) ) for chunk in stream_resp: token = chunk.choices[0].delta.content if token: yield f"data: {json.dumps({'type': 'token', 'text': token})}\n\n" await asyncio.sleep(0) yield "data: {\"type\": \"done\"}\n\n" return StreamingResponse(generate(), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}) @app.post("/compare") def compare(req: QueryRequest): """So sánh Text RAG vs ColPali retrieval trên cùng câu hỏi.""" from src.colpali_retriever import query as colpali_query _, text_sources, _ = retrieve(req.question, top_k=req.top_k) try: colpali_hits = colpali_query(req.question, top_k=req.top_k) except Exception as e: colpali_hits = [] return {"text_rag": text_sources, "colpali": colpali_hits} @app.get("/health") def health(): return {"status": "ok"}