thaidinhz1 commited on
Commit
b18c987
·
1 Parent(s): 6a99458

feat: streaming response via SSE with token-by-token display

Browse files
Files changed (3) hide show
  1. api.py +27 -2
  2. src/rag.py +19 -14
  3. 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 answer(question: str, top_k: int = 5) -> dict:
35
- """RAG pipeline: rewrite -> embed -> hybrid retrieve -> rerank -> generate."""
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
- all_candidates = candidates + uploaded
44
- contexts = rerank(question, all_candidates, top_k=top_k)
 
 
 
 
 
 
 
 
 
 
 
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('/query', {
285
  method: 'POST',
286
  headers: { 'Content-Type': 'application/json' },
287
  body: JSON.stringify({ question: q, top_k: 5 })
288
  });
289
- const data = await res.json();
290
- item.querySelector('.bubble-a').innerHTML =
291
- `<div class="answer-text">${highlightCitations(data.answer)}</div>${renderSources(data.sources)}`;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
  }