Spaces:
Sleeping
Sleeping
Commit ·
b18c987
1
Parent(s): 6a99458
feat: streaming response via SSE with token-by-token display
Browse files- api.py +27 -2
- src/rag.py +19 -14
- static/index.html +32 -4
api.py
CHANGED
|
@@ -1,9 +1,10 @@
|
|
|
|
|
| 1 |
from contextlib import asynccontextmanager
|
| 2 |
from fastapi import FastAPI, UploadFile, File, HTTPException
|
| 3 |
from fastapi.staticfiles import StaticFiles
|
| 4 |
-
from fastapi.responses import HTMLResponse
|
| 5 |
from pydantic import BaseModel
|
| 6 |
-
from src.rag import answer
|
| 7 |
from src.reranker import _get_model
|
| 8 |
from src.uploader import ingest_pdf, get_uploaded_files
|
| 9 |
|
|
@@ -68,6 +69,30 @@ def documents():
|
|
| 68 |
return get_uploaded_files()
|
| 69 |
|
| 70 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 71 |
@app.get("/health")
|
| 72 |
def health():
|
| 73 |
return {"status": "ok"}
|
|
|
|
| 1 |
+
import json
|
| 2 |
from contextlib import asynccontextmanager
|
| 3 |
from fastapi import FastAPI, UploadFile, File, HTTPException
|
| 4 |
from fastapi.staticfiles import StaticFiles
|
| 5 |
+
from fastapi.responses import HTMLResponse, StreamingResponse
|
| 6 |
from pydantic import BaseModel
|
| 7 |
+
from src.rag import answer, retrieve
|
| 8 |
from src.reranker import _get_model
|
| 9 |
from src.uploader import ingest_pdf, get_uploaded_files
|
| 10 |
|
|
|
|
| 69 |
return get_uploaded_files()
|
| 70 |
|
| 71 |
|
| 72 |
+
@app.post("/stream")
|
| 73 |
+
async def stream(req: QueryRequest):
|
| 74 |
+
def generate():
|
| 75 |
+
contexts, sources = retrieve(req.question, top_k=req.top_k)
|
| 76 |
+
# Gửi sources trước
|
| 77 |
+
yield f"data: {json.dumps({'type': 'sources', 'sources': sources})}\n\n"
|
| 78 |
+
# Stream từng token
|
| 79 |
+
from src.rag import build_prompt, client, LLM_MODEL
|
| 80 |
+
prompt = build_prompt(req.question, contexts)
|
| 81 |
+
stream_resp = client.chat.completions.create(
|
| 82 |
+
model=LLM_MODEL,
|
| 83 |
+
messages=[{"role": "user", "content": prompt}],
|
| 84 |
+
stream=True,
|
| 85 |
+
)
|
| 86 |
+
for chunk in stream_resp:
|
| 87 |
+
token = chunk.choices[0].delta.content
|
| 88 |
+
if token:
|
| 89 |
+
yield f"data: {json.dumps({'type': 'token', 'text': token})}\n\n"
|
| 90 |
+
yield "data: {\"type\": \"done\"}\n\n"
|
| 91 |
+
|
| 92 |
+
return StreamingResponse(generate(), media_type="text/event-stream",
|
| 93 |
+
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"})
|
| 94 |
+
|
| 95 |
+
|
| 96 |
@app.get("/health")
|
| 97 |
def health():
|
| 98 |
return {"status": "ok"}
|
src/rag.py
CHANGED
|
@@ -31,8 +31,8 @@ CÂU HỎI: {question}
|
|
| 31 |
TRẢ LỜI:"""
|
| 32 |
|
| 33 |
|
| 34 |
-
def
|
| 35 |
-
"""
|
| 36 |
rewritten = rewrite_query(question)
|
| 37 |
if rewritten != question:
|
| 38 |
print(f"Query rewritten: {rewritten}")
|
|
@@ -40,25 +40,30 @@ def answer(question: str, top_k: int = 5) -> dict:
|
|
| 40 |
query_vec = embed_query(rewritten)
|
| 41 |
candidates = vector_query(query_vec, query_text=rewritten, top_k=top_k * 3)
|
| 42 |
uploaded = search_uploaded(query_vec, rewritten, top_k=top_k * 2)
|
| 43 |
-
|
| 44 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
|
|
|
|
|
|
|
|
|
|
| 46 |
prompt = build_prompt(question, contexts)
|
| 47 |
response = client.chat.completions.create(
|
| 48 |
model=LLM_MODEL,
|
| 49 |
messages=[{"role": "user", "content": prompt}],
|
| 50 |
)
|
| 51 |
-
|
| 52 |
return {
|
| 53 |
"answer": response.choices[0].message.content,
|
| 54 |
"contexts": [c["text"] for c in contexts],
|
| 55 |
-
"sources":
|
| 56 |
-
{
|
| 57 |
-
"source": c["metadata"]["source"],
|
| 58 |
-
"page": c["metadata"]["page"],
|
| 59 |
-
"rrf_score": round(c["score"], 3),
|
| 60 |
-
"rerank_score": round(c["rerank_score"], 3),
|
| 61 |
-
}
|
| 62 |
-
for c in contexts
|
| 63 |
-
],
|
| 64 |
}
|
|
|
|
| 31 |
TRẢ LỜI:"""
|
| 32 |
|
| 33 |
|
| 34 |
+
def retrieve(question: str, top_k: int = 5):
|
| 35 |
+
"""Rewrite -> embed -> hybrid retrieve -> rerank. Returns (contexts, sources)."""
|
| 36 |
rewritten = rewrite_query(question)
|
| 37 |
if rewritten != question:
|
| 38 |
print(f"Query rewritten: {rewritten}")
|
|
|
|
| 40 |
query_vec = embed_query(rewritten)
|
| 41 |
candidates = vector_query(query_vec, query_text=rewritten, top_k=top_k * 3)
|
| 42 |
uploaded = search_uploaded(query_vec, rewritten, top_k=top_k * 2)
|
| 43 |
+
contexts = rerank(question, candidates + uploaded, top_k=top_k)
|
| 44 |
+
|
| 45 |
+
sources = [
|
| 46 |
+
{
|
| 47 |
+
"source": c["metadata"]["source"],
|
| 48 |
+
"page": c["metadata"]["page"],
|
| 49 |
+
"rrf_score": round(c["score"], 3),
|
| 50 |
+
"rerank_score": round(c["rerank_score"], 3),
|
| 51 |
+
}
|
| 52 |
+
for c in contexts
|
| 53 |
+
]
|
| 54 |
+
return contexts, sources
|
| 55 |
+
|
| 56 |
|
| 57 |
+
def answer(question: str, top_k: int = 5) -> dict:
|
| 58 |
+
"""RAG pipeline: rewrite -> embed -> hybrid retrieve -> rerank -> generate."""
|
| 59 |
+
contexts, sources = retrieve(question, top_k)
|
| 60 |
prompt = build_prompt(question, contexts)
|
| 61 |
response = client.chat.completions.create(
|
| 62 |
model=LLM_MODEL,
|
| 63 |
messages=[{"role": "user", "content": prompt}],
|
| 64 |
)
|
|
|
|
| 65 |
return {
|
| 66 |
"answer": response.choices[0].message.content,
|
| 67 |
"contexts": [c["text"] for c in contexts],
|
| 68 |
+
"sources": sources,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 69 |
}
|
static/index.html
CHANGED
|
@@ -281,14 +281,42 @@
|
|
| 281 |
btn.disabled = true;
|
| 282 |
|
| 283 |
try {
|
| 284 |
-
const res = await fetch('/
|
| 285 |
method: 'POST',
|
| 286 |
headers: { 'Content-Type': 'application/json' },
|
| 287 |
body: JSON.stringify({ question: q, top_k: 5 })
|
| 288 |
});
|
| 289 |
-
|
| 290 |
-
item.querySelector('.bubble-a')
|
| 291 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 292 |
} catch (e) {
|
| 293 |
item.querySelector('.bubble-a').innerHTML = `<span class="error-text">Lỗi: ${e.message}</span>`;
|
| 294 |
}
|
|
|
|
| 281 |
btn.disabled = true;
|
| 282 |
|
| 283 |
try {
|
| 284 |
+
const res = await fetch('/stream', {
|
| 285 |
method: 'POST',
|
| 286 |
headers: { 'Content-Type': 'application/json' },
|
| 287 |
body: JSON.stringify({ question: q, top_k: 5 })
|
| 288 |
});
|
| 289 |
+
|
| 290 |
+
const bubble = item.querySelector('.bubble-a');
|
| 291 |
+
bubble.innerHTML = '<div class="answer-text"></div>';
|
| 292 |
+
const answerEl = bubble.querySelector('.answer-text');
|
| 293 |
+
let fullText = '';
|
| 294 |
+
let sources = [];
|
| 295 |
+
|
| 296 |
+
const reader = res.body.getReader();
|
| 297 |
+
const decoder = new TextDecoder();
|
| 298 |
+
let buffer = '';
|
| 299 |
+
|
| 300 |
+
while (true) {
|
| 301 |
+
const { done, value } = await reader.read();
|
| 302 |
+
if (done) break;
|
| 303 |
+
buffer += decoder.decode(value, { stream: true });
|
| 304 |
+
const lines = buffer.split('\n');
|
| 305 |
+
buffer = lines.pop();
|
| 306 |
+
for (const line of lines) {
|
| 307 |
+
if (!line.startsWith('data: ')) continue;
|
| 308 |
+
const msg = JSON.parse(line.slice(6));
|
| 309 |
+
if (msg.type === 'sources') {
|
| 310 |
+
sources = msg.sources;
|
| 311 |
+
} else if (msg.type === 'token') {
|
| 312 |
+
fullText += msg.text;
|
| 313 |
+
answerEl.innerHTML = highlightCitations(fullText);
|
| 314 |
+
item.scrollIntoView({ behavior: 'smooth', block: 'end' });
|
| 315 |
+
} else if (msg.type === 'done') {
|
| 316 |
+
bubble.innerHTML = `<div class="answer-text">${highlightCitations(fullText)}</div>${renderSources(sources)}`;
|
| 317 |
+
}
|
| 318 |
+
}
|
| 319 |
+
}
|
| 320 |
} catch (e) {
|
| 321 |
item.querySelector('.bubble-a').innerHTML = `<span class="error-text">Lỗi: ${e.message}</span>`;
|
| 322 |
}
|