rag-vietnamese / src /rag.py
thaidinhz1's picture
feat: streaming response via SSE with token-by-token display
b18c987
Raw History Blame
2.41 kB
import os
from groq import Groq
from dotenv import load_dotenv
from src.embedder import embed_query
from src.vector_store import query as vector_query
from src.reranker import rerank
from src.query_rewriter import rewrite_query
from src.uploader import search_uploaded
load_dotenv()
client = Groq(api_key=os.environ["GROQ_API_KEY"].strip())
LLM_MODEL = "llama-3.3-70b-versatile"
def build_prompt(question: str, contexts: list[dict]) -> str:
context_text = ""
for i, ctx in enumerate(contexts, 1):
meta = ctx["metadata"]
context_text += f"[{i}] (File: {meta['source']}, Trang: {meta['page']})\n{ctx['text']}\n\n"
return f"""Bạn là trợ lý trả lời câu hỏi dựa trên tài liệu được cung cấp.
Chỉ dùng thông tin trong phần NGỮ CẢNH bên dưới để trả lời.
Trích dẫn nguồn bằng số [1], [2],... tương ứng với từng đoạn bạn sử dụng.
Nếu không tìm thấy thông tin, hãy nói "Không tìm thấy thông tin trong tài liệu."
NGỮ CẢNH:
{context_text}
CÂU HỎI: {question}
TRẢ LỜI:"""
def retrieve(question: str, top_k: int = 5):
"""Rewrite -> embed -> hybrid retrieve -> rerank. Returns (contexts, sources)."""
rewritten = rewrite_query(question)
if rewritten != question:
print(f"Query rewritten: {rewritten}")
query_vec = embed_query(rewritten)
candidates = vector_query(query_vec, query_text=rewritten, top_k=top_k * 3)
uploaded = search_uploaded(query_vec, rewritten, top_k=top_k * 2)
contexts = rerank(question, candidates + uploaded, top_k=top_k)
sources = [
{
"source": c["metadata"]["source"],
"page": c["metadata"]["page"],
"rrf_score": round(c["score"], 3),
"rerank_score": round(c["rerank_score"], 3),
}
for c in contexts
]
return contexts, sources
def answer(question: str, top_k: int = 5) -> dict:
"""RAG pipeline: rewrite -> embed -> hybrid retrieve -> rerank -> generate."""
contexts, sources = retrieve(question, top_k)
prompt = build_prompt(question, contexts)
response = client.chat.completions.create(
model=LLM_MODEL,
messages=[{"role": "user", "content": prompt}],
)
return {
"answer": response.choices[0].message.content,
"contexts": [c["text"] for c in contexts],
"sources": sources,
}