rag-document-qa / generation /generator.py
Amrita P
minor fixes-2
57de4ca
Raw
History Blame Contribute Delete
7.86 kB
import os
import logging
from dataclasses import dataclass, field
from typing import Iterator
from groq import Groq
from dotenv import load_dotenv
from retrieval.index import SearchResult
load_dotenv()
logger = logging.getLogger(__name__)
_MODEL_NAME = "llama-3.3-70b-versatile"
_ENV_KEY = "GROQ_API_KEY"
_SYSTEM_PROMPT = """\
You are a helpful document Q&A assistant. Follow these rules:
1. Answer the question using the information in the numbered context sources below.
2. Cite your claims inline with [Source N] where N is the source number.
3. If multiple sources support a claim, cite all: [Source 1][Source 3].
4. Synthesize information across multiple sources to give a complete answer.
5. If context is partial or fragmented, do your best to piece together a coherent answer from what is available.
6. Only say you cannot answer if the context is genuinely about a completely different topic.
7. Give detailed, thorough answers - not one-line summaries.\
"""
_LOW_CONFIDENCE_MESSAGE = (
"I don't have enough information in the uploaded documents to answer this question. "
"Try uploading a more relevant document or rephrasing your question."
)
_MEDIUM_CONFIDENCE_NOTE = (
"\n\n⚠️ **Low confidence:** The retrieved context may not fully address your question."
)
@dataclass
class Answer:
question: str
answer: str
sources: list[dict] = field(default_factory=list)
confidence_score: float = 0.0
confidence_level: str = "high" # "high" | "medium" | "low"
class Generator:
"""Wraps Groq (llama-3.3-70b-versatile) for grounded, citation-aware answer generation.
The API key is read from the GROQ_API_KEY environment variable (or a
.env file in the working directory). A low temperature keeps answers
factual; max_tokens caps runaway responses.
"""
def __init__(
self,
model_name: str = _MODEL_NAME,
low_confidence_threshold: float = 0.3,
medium_confidence_threshold: float = 0.5,
) -> None:
api_key = os.getenv(_ENV_KEY)
if not api_key:
raise EnvironmentError(
f"Groq API key not found. Set the {_ENV_KEY} environment variable "
"or add it to a .env file."
)
self._low_threshold = low_confidence_threshold
self._medium_threshold = medium_confidence_threshold
self._model_name = model_name
self._client = Groq(api_key=api_key)
logger.info("Generator initialised with model: %s", model_name)
def _confidence_level(self, max_score: float) -> str:
if max_score < self._low_threshold:
return "low"
if max_score < self._medium_threshold:
return "medium"
return "high"
def generate_answer(
self,
question: str,
results: list[SearchResult],
max_score: float = 1.0,
) -> Answer:
"""Build a grounded prompt from retrieved chunks and call Groq.
Args:
question: The user's question string.
results: Ranked SearchResult list from VectorIndex.search().
Each result's metadata must contain 'text', 'source',
and 'page_num' keys (set by IngestionPipeline).
max_score: Highest retrieval score from the search step, used to
gate the API call and set confidence metadata.
Returns:
Answer with generated text, source metadata, and confidence fields.
When max_score is below the low threshold, the LLM is not called.
"""
if not results:
return Answer(
question=question,
answer="No relevant context was retrieved to answer this question.",
sources=[],
confidence_score=0.0,
confidence_level="low",
)
level = self._confidence_level(max_score)
if level == "low":
logger.debug("Skipping LLM call: max_score=%.4f below low threshold %.2f", max_score, self._low_threshold)
return Answer(
question=question,
answer=_LOW_CONFIDENCE_MESSAGE,
sources=[r.metadata for r in results],
confidence_score=max_score,
confidence_level="low",
)
context_block = _build_context(results)
prompt = (
f"CONTEXT:\n{context_block}\n\n"
f"QUESTION:\n{question}\n\n"
"ANSWER:"
)
logger.debug("Sending prompt to Groq (%d chars)", len(prompt))
response = self._client.chat.completions.create(
model=self._model_name,
messages=[
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": prompt},
],
temperature=0.1,
max_tokens=2048,
)
answer_text = response.choices[0].message.content.strip()
if level == "medium":
answer_text += _MEDIUM_CONFIDENCE_NOTE
return Answer(
question=question,
answer=answer_text,
sources=[r.metadata for r in results],
confidence_score=max_score,
confidence_level=level,
)
def generate_answer_stream(
self,
question: str,
results: list[SearchResult],
max_score: float = 1.0,
) -> Iterator[str]:
"""Stream the answer token-by-token using Groq's streaming API.
Args:
question: The user's question string.
results: Ranked SearchResult list (same as generate_answer).
max_score: Highest retrieval score; gates the API call and
appends a warning note for medium-confidence answers.
Yields:
Text chunks as they arrive from Groq.
Raises:
Exception: Re-raises any Groq API error so the caller can handle it.
"""
if not results:
yield "No relevant context was retrieved to answer this question."
return
level = self._confidence_level(max_score)
if level == "low":
logger.debug("Skipping LLM call: max_score=%.4f below low threshold %.2f", max_score, self._low_threshold)
yield _LOW_CONFIDENCE_MESSAGE
return
context_block = _build_context(results)
prompt = (
f"CONTEXT:\n{context_block}\n\n"
f"QUESTION:\n{question}\n\n"
"ANSWER:"
)
logger.debug("Sending streaming prompt to Groq (%d chars)", len(prompt))
stream = self._client.chat.completions.create(
model=self._model_name,
messages=[
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": prompt},
],
temperature=0.1,
max_tokens=2048,
stream=True,
)
for chunk in stream:
if chunk.choices[0].delta.content:
yield chunk.choices[0].delta.content
if level == "medium":
yield _MEDIUM_CONFIDENCE_NOTE
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _build_context(results: list[SearchResult]) -> str:
"""Format retrieved chunks into a numbered source block for the prompt."""
parts: list[str] = []
for i, result in enumerate(results, start=1):
meta = result.metadata
header = (
f"[Source {i}] "
f"{meta.get('source', 'unknown')}, "
f"page {meta.get('page_num', '?')}"
)
text = meta.get("text", "").strip()
parts.append(f"{header}\n{text}")
return "\n\n".join(parts)