rag-vietnamese / api.py
thaidinhz1's picture
Remove upload feature: broken with scanned PDFs, replaced by fixed corpus
0918009
Raw History Blame Contribute Delete
3.23 kB
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"}