Spaces:
Sleeping
Sleeping
Download api.py from thaidinhz1/rag-vietnamese: direct link, hf CLI and curl.
- Browser
- Download file 3.23 kB
-
https://huggingface.co/spaces/thaidinhz1/rag-vietnamese/resolve/main/api.py
- Command line
-
hf download hf://spaces/thaidinhz1/rag-vietnamese/api.py
-
curl -L -o api.py https://huggingface.co/spaces/thaidinhz1/rag-vietnamese/resolve/main/api.py
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 | |
| 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 = "" | |
| def index(): | |
| with open("static/index.html", encoding="utf-8") as f: | |
| return f.read() | |
| 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", ""), | |
| ) | |
| 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"}) | |
| 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} | |
| def health(): | |
| return {"status": "ok"} |