"""
Contextual Retrieval — sinh 1-2 câu context cho mỗi chunk dựa trên toàn văn tài liệu.
Context được prepend vào chunk trước khi embed + BM25, giúp retrieval chính xác hơn.
Tham khảo: Anthropic Contextual Retrieval paper (~35% giảm retrieval failure).
"""
import os
from groq import Groq
from dotenv import load_dotenv
load_dotenv()
_client = Groq(api_key=os.environ["GROQ_API_KEY"].strip())
CONTEXT_MODEL = "llama-3.1-8b-instant"
CONTEXT_PROMPT = """Dưới đây là toàn bộ tài liệu:
{doc_text}
Đây là đoạn văn cần đặt vào ngữ cảnh:
{chunk_text}
Hãy viết 1-2 câu ngắn gọn mô tả vị trí và vai trò của đoạn này trong tài liệu, \
giúp người đọc hiểu đoạn này nằm ở đâu và nói về điều gì.
Chỉ trả về 1-2 câu đó, không giải thích thêm."""
def generate_context(doc_text: str, chunk_text: str) -> str:
"""Sinh context cho 1 chunk. Trả về context string."""
# Giới hạn doc_text để tránh vượt context window
max_doc_chars = 3000
if len(doc_text) > max_doc_chars:
doc_text = doc_text[:max_doc_chars] + "\n...[nội dung tiếp theo]..."
import time
prompt = CONTEXT_PROMPT.format(doc_text=doc_text, chunk_text=chunk_text)
for attempt in range(5):
try:
response = _client.chat.completions.create(
model=CONTEXT_MODEL,
messages=[{"role": "user", "content": prompt}],
max_tokens=150,
temperature=0.1,
)
return response.choices[0].message.content.strip()
except Exception as e:
msg = str(e)
if "429" in msg:
if "per_day" in msg.lower() or "TPD" in msg or "limit 100000" in msg:
print(f" [contextualizer] hết quota ngày, dừng sinh context")
raise RuntimeError("QUOTA_DAY_EXHAUSTED") from e
wait = 60 * (attempt + 1)
print(f" [contextualizer] rate limit, chờ {wait}s rồi thử lại...")
time.sleep(wait)
else:
print(f" [contextualizer] lỗi: {e}, bỏ qua context")
return ""
return ""
def add_context_to_chunks(chunks: list[dict], source_pages: list[dict]) -> list[dict]:
"""
Prepend context vào text của mỗi chunk.
source_pages: list các page dict có 'text' và 'metadata.source' để build doc_text.
"""
# Group pages theo source file
from collections import defaultdict
doc_texts = defaultdict(list)
for page in source_pages:
src = page["metadata"].get("source", "unknown")
doc_texts[src].append(page["text"])
# Full text per document
full_docs = {src: "\n\n".join(pages) for src, pages in doc_texts.items()}
import time
enriched = []
total = len(chunks)
for i, chunk in enumerate(chunks):
src = chunk["metadata"].get("source", "unknown")
doc_text = full_docs.get(src, "")
print(f" Generating context {i+1}/{total}: {src}...")
try:
context = generate_context(doc_text, chunk["text"])
except RuntimeError as e:
if "QUOTA_DAY_EXHAUSTED" in str(e):
print(f" [contextualizer] quota ngày hết tại chunk {i+1}/{total}, dừng lại")
# Append remaining chunks without context
for remaining in chunks[i:]:
r = dict(remaining)
r["metadata"] = {**remaining["metadata"], "has_context": False}
enriched.append(r)
return enriched
context = ""
time.sleep(4) # Gemini Flash free tier: 15 RPM
new_chunk = dict(chunk)
if context:
new_chunk["text"] = f"{context}\n\n{chunk['text']}"
new_chunk["metadata"] = {**chunk["metadata"], "has_context": True}
else:
new_chunk["metadata"] = {**chunk["metadata"], "has_context": False}
enriched.append(new_chunk)
return enriched