rag-uigreenmetric / src /pipeline.py
fortunius's picture
[FIX] correct default budget to 70k in src/pipeline.py
2b5e095 verified
Raw
History Blame Contribute Delete
6.53 kB
"""Pipeline orchestrator for the UI GreenMetric RAG system.
Wires the router, retriever, RAG Fusion paraphrasing, and generator
into a single end-to-end ``ask()`` entry point.
"""
import os
import time
import json
from src.router import route, paraphrase
from src.retriever import retrieve, retrieve_multi, aggregate_stats
from src.generator import generate
from src.budget import BudgetManager, MemoryBudgetStore, HFBudgetStore
from src.conversation import log_conversation, get_logger
def _flush_logs() -> None:
logger = get_logger()
if logger:
logger.flush()
_RERANK_ENABLED = os.getenv("RAG_RERANK", "1") == "1"
_BUDGET_BLOCKED_RESPONSE = {
"answer": "Daily token budget reached. Please try again tomorrow.",
"route": {"source": "none", "csv_source": None, "query_type": "lookup"},
"retrieved": 0,
"contexts": [],
"low_confidence": False,
"rerank_ms": 0.0,
}
# Budget: use HF dataset store if repo configured, else in-memory
_budget_repo = os.getenv("RAG_BUDGET_REPO", "")
_budget = BudgetManager(
store=HFBudgetStore(_budget_repo) if _budget_repo else MemoryBudgetStore(),
daily_cap=int(os.getenv("RAG_BUDGET_TOKENS", "70000")),
)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def ask(
query: str,
*,
_route_result: dict | None = None,
_fusion_queries: list[str] | None = None,
_on_status: "Callable[[str], None] | None" = None,
) -> dict:
"""Run the full RAG pipeline.
Flow:
source == "none"
→ immediate polite refusal (no LLM / embedding calls).
query_type == "aggregate"
→ fetch all chunks via exact metadata match;
skip paraphrase and reranker.
everything else
→ paraphrase (3 variants) → multi-query retrieval
→ RRF merge → (optional reranker) → top 7.
Parameters:
query: The user's question.
Returns:
dict with keys:
* ``"answer"`` — the generated answer string.
* ``"route"`` — the dict returned by :func:`router.route`.
* ``"retrieved"`` — the number of chunks after retrieval.
* ``"contexts"`` — list of retrieved chunk content strings.
* ``"low_confidence"`` — ``True`` when the top chunk exceeded
the 0.6 cosine-distance warning threshold.
* ``"rerank_ms"`` — milliseconds spent in the reranker
(0.0 when skipped).
"""
# ── budget guard ─────────────────────────────────────────────────
if _budget.exceeded():
return _BUDGET_BLOCKED_RESPONSE
if _route_result is not None:
route_result = _route_result # pre-computed, no LLM call
else:
if _on_status:
_on_status("Routing query...")
route_result, route_tokens = route(query)
_budget.track(route_tokens)
# ── none ────────────────────────────────────────────────────────────
if route_result["source"] == "none":
return {
"answer": "I don't have the required information to answer this question.",
"route": route_result,
"retrieved": 0,
"contexts": [],
"low_confidence": False,
"rerank_ms": 0.0,
}
# ── aggregate ───────────────────────────────────────────────────────
if route_result["query_type"] == "aggregate":
if _on_status:
_on_status("Retrieving context...")
agg_source = route_result.get("csv_source") or route_result["source"]
stats = aggregate_stats(agg_source)
if stats:
context = [{"content": stats, "distance": 0.0,
"metadata": {"source": "aggregator", "chunk_type": "stats"}}]
else:
context = retrieve(query, route_result) # fallback to _fetch_all
if _budget.exceeded():
return _BUDGET_BLOCKED_RESPONSE
if _on_status:
_on_status("Generating answer...")
answer, gen_tokens = generate(query, context, query_type="aggregate")
_budget.track(gen_tokens)
log_conversation(query, answer, route_result,
[c["content"] for c in context], gen_tokens)
_flush_logs()
return {
"answer": answer,
"route": route_result,
"retrieved": len(context),
"contexts": [c["content"] for c in context],
"low_confidence": False,
"rerank_ms": 0.0,
}
# ── lookup / both → RAG Fusion ──────────────────────────────────────
if _fusion_queries is not None:
queries = _fusion_queries
else:
if _on_status:
_on_status("Paraphrasing query...")
variants, para_tokens = paraphrase(query)
_budget.track(para_tokens)
queries = [query] + variants
if _on_status:
_on_status("Retrieving context...")
context = retrieve_multi(queries, route_result)
rerank_ms = 0.0
if _RERANK_ENABLED and context:
if _on_status:
_on_status("Reranking results...")
from src.reranker import rerank
t0 = time.perf_counter()
context = rerank(query, context, top_n=7)
rerank_ms = (time.perf_counter() - t0) * 1000
else:
context = context[:7]
if _budget.exceeded():
return _BUDGET_BLOCKED_RESPONSE
if _on_status:
_on_status("Generating answer...")
answer, gen_tokens = generate(query, context, query_type=route_result["query_type"])
_budget.track(gen_tokens)
low_confidence = (
bool(context)
and context[0].get("distance", 0.0) > 0.6
)
log_conversation(query, answer, route_result,
[c["content"] for c in context], gen_tokens)
_flush_logs()
return {
"answer": answer,
"route": route_result,
"retrieved": len(context),
"contexts": [c["content"] for c in context],
"low_confidence": low_confidence,
"rerank_ms": rerank_ms,
}