agentic-rag / backend /app /api /query_router.py
vksepm
fixed plotly visualization and added evals
73af3c0
Raw
History Blame Contribute Delete
9.42 kB
"""
Query endpoint β€” the main pipeline entry point.
POST /api/v1/query β€” run an agentic RAG query (authenticated)
GET /api/v1/query/history β€” paginated query history for the current user
"""
import asyncio
import json
import logging
from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from concurrent.futures import ThreadPoolExecutor
from app.agent.agent import run_agent
from app.auth import get_current_user
from app.config import get_settings
from app.database import get_db
from app.evaluation.trulens_eval import EvalScores, evaluate_async
from app.models import EvaluationResult, QueryLog, User
from app.visualization.viz_agent import VizCacheEntry, VIZ_CACHE, run_viz_agent
_viz_executor = ThreadPoolExecutor(max_workers=4, thread_name_prefix="viz-agent")
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/query", tags=["query"])
# ── Pydantic schemas ──────────────────────────────────────────────────────────
class QueryRequest(BaseModel):
query: str
include_visualization: bool = False
class Citation(BaseModel):
index: int
title: str
source_url: str
full_citation: str
class QueryResponse(BaseModel):
id: str
query: str
answer: str
citations: list[Citation]
chart_data: dict | None
model_provider: str
agent_steps: int
created_at: datetime
class QueryHistoryItem(BaseModel):
id: str
query_text: str
response_text: str | None
model_provider: str | None
agent_steps: int | None
created_at: datetime
class EvalScoresPayload(BaseModel):
relevance_score: float | None
groundedness_score: float | None
answer_relevance_score: float | None
class EvalStatusResponse(BaseModel):
status: str # "pending" | "complete"
scores: EvalScoresPayload | None = None
# ── Helper ─────────────────────────────────────────────────────────────────────
def _persist_evaluation(
eval_scores: EvalScores,
query_log_id: str,
loop: asyncio.AbstractEventLoop,
db_session_factory,
) -> None:
"""
Background callback: write TruLens scores to evaluation_results table.
Runs when the evaluation future resolves.
"""
async def _write():
async with db_session_factory() as session:
row = EvaluationResult(
query_log_id=query_log_id,
relevance_score=eval_scores.relevance_score,
groundedness_score=eval_scores.groundedness_score,
answer_relevance_score=eval_scores.answer_relevance_score,
trulens_record_id=eval_scores.trulens_record_id,
)
session.add(row)
await session.commit()
asyncio.run_coroutine_threadsafe(_write(), loop)
# ── Endpoints ─────────────────────────────────────────────────────────────────
@router.post("", response_model=QueryResponse)
async def submit_query(
body: QueryRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
settings = get_settings()
if not body.query.strip():
raise HTTPException(status_code=422, detail="Query must not be empty")
# ── Run the agent (synchronous work dispatched to thread pool) ────────────
loop = asyncio.get_event_loop()
agent_result: dict = await loop.run_in_executor(
None, run_agent, current_user.id, body.query, settings
)
answer: str = agent_result.get("answer", "")
citations_raw: list = agent_result.get("citations", [])
chart_data: dict | None = agent_result.get("chart_data")
agent_steps: int = agent_result.get("agent_steps", 0)
# ── Persist query log ─────────────────────────────────────────────────────
log_entry = QueryLog(
user_id=current_user.id,
query_text=body.query,
response_text=answer,
model_provider=settings.model_provider,
agent_steps=agent_steps,
)
db.add(log_entry)
await db.commit()
await db.refresh(log_entry)
# ── Schedule TruLens evaluation asynchronously ────────────────────────────
# Prefer authoritative pgvector_retriever output (reranked top-5 + chunk text);
# JSON "context" may be truncated; citation-only fallback lacks Content:.
citation_fallback = "\n\n".join(
c.get("full_citation", "") + "\n" + c.get("title", "")
for c in citations_raw
)
context_text = (
agent_result.get("reranked_context", "")
or agent_result.get("context", "")
or citation_fallback
)
if not agent_result.get("reranked_context") and citations_raw:
logger.debug(
"TruLens context: empty reranked_context but citations present β€” using JSON context or citation fallback"
)
eval_future = await evaluate_async(
query=body.query,
context=context_text,
answer=answer,
query_log_id=log_entry.id,
settings=settings,
)
# Wire up persistence callback when evaluation completes (fire-and-forget)
from app.database import AsyncSessionLocal
def _on_eval_done(future):
try:
scores: EvalScores = future.result()
_persist_evaluation(scores, log_entry.id, loop, AsyncSessionLocal)
except Exception as exc:
logger.error("Evaluation persistence failed: %s", exc)
eval_future.add_done_callback(_on_eval_done)
# Normalise citations into typed objects
citations = [
Citation(
index=c.get("index", i + 1),
title=c.get("title", ""),
source_url=c.get("source_url", ""),
full_citation=c.get("full_citation", ""),
)
for i, c in enumerate(citations_raw)
]
# ── Fire visualization agent (fire-and-forget, same pattern as TruLens) ────
if body.include_visualization:
# Seed a "pending" entry immediately so the poll endpoint returns 200
VIZ_CACHE[log_entry.id] = VizCacheEntry(status="pending")
_viz_executor.submit(run_viz_agent, log_entry.id, body.query, answer, settings)
logger.debug("Visualization agent submitted for query_id=%s", log_entry.id)
return QueryResponse(
id=log_entry.id,
query=body.query,
answer=answer,
citations=citations,
chart_data=chart_data,
model_provider=settings.model_provider,
agent_steps=agent_steps,
created_at=log_entry.created_at,
)
@router.get("/{query_id}/evaluation", response_model=EvalStatusResponse)
async def get_evaluation(
query_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""
Poll for TruLens RAG Triad scores for a completed query.
Returns {"status": "pending"} while the background eval thread is still running,
and {"status": "complete", "scores": {...}} once the row is persisted.
"""
# Verify query exists and belongs to this user
log_result = await db.execute(
select(QueryLog).where(
QueryLog.id == query_id,
QueryLog.user_id == current_user.id,
)
)
if log_result.scalar_one_or_none() is None:
raise HTTPException(status_code=404, detail="Query not found")
eval_result = await db.execute(
select(EvaluationResult).where(EvaluationResult.query_log_id == query_id)
)
row = eval_result.scalar_one_or_none()
if row is None:
return EvalStatusResponse(status="pending")
return EvalStatusResponse(
status="complete",
scores=EvalScoresPayload(
relevance_score=float(row.relevance_score) if row.relevance_score is not None else None,
groundedness_score=float(row.groundedness_score) if row.groundedness_score is not None else None,
answer_relevance_score=float(row.answer_relevance_score) if row.answer_relevance_score is not None else None,
),
)
@router.get("/history", response_model=list[QueryHistoryItem])
async def query_history(
limit: int = 20,
offset: int = 0,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(QueryLog)
.where(QueryLog.user_id == current_user.id)
.order_by(QueryLog.created_at.desc())
.limit(limit)
.offset(offset)
)
logs = result.scalars().all()
return [
QueryHistoryItem(
id=log.id,
query_text=log.query_text,
response_text=log.response_text,
model_provider=log.model_provider,
agent_steps=log.agent_steps,
created_at=log.created_at,
)
for log in logs
]