Ukweli / app /api /agent.py
Benard John
feat: implement full-stack architecture with database models, authentication services, and web dashboard components
37b5223
Raw
History Blame Contribute Delete
6.36 kB
"""
Ukweli — Agent Streaming Endpoint
POST /agent/stream — Server-Sent Events (SSE) streaming for real-time RAG responses.
Architecture Section 4.2 — Agent Endpoint (Streaming).
"""
from __future__ import annotations
import json
import logging
import time
import uuid
from collections.abc import AsyncGenerator
from fastapi import APIRouter, Depends
from sse_starlette.sse import EventSourceResponse
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.auth import AuthContext, require_auth
from app.db.session import get_db_session
from app.dependencies import get_citation_formatter, get_llm_gateway, get_retrieval_orchestrator
from app.models.database import QueryLog
from app.models.schemas import AgentStreamRequest
from app.services.citation.formatter import CitationFormatter
from app.services.llm.gateway import LLMGateway
from app.services.llm.guardrails import check_query_safety
from app.services.llm.prompts import build_rag_prompt
from app.services.retrieval.orchestrator import RetrievalOrchestrator
logger = logging.getLogger("ukweli.api.agent")
router = APIRouter(tags=["Agent"])
@router.post("/agent/stream")
async def agent_stream(
request: AgentStreamRequest,
db: AsyncSession = Depends(get_db_session),
auth: AuthContext = Depends(require_auth),
retriever: RetrievalOrchestrator = Depends(get_retrieval_orchestrator),
llm: LLMGateway = Depends(get_llm_gateway),
citation_fmt: CitationFormatter = Depends(get_citation_formatter),
):
"""
SSE streaming endpoint for agent integrations (WhatsApp, mobile, etc.).
Event types emitted:
- retrieval: search status and source count
- citation: individual citation as discovered
- delta: incremental answer text
- done: final payload with all citations and confidence
- error: if something goes wrong
"""
async def event_generator() -> AsyncGenerator[dict, None]:
start_time = time.monotonic()
query_id = uuid.uuid4()
# Safety check
safety = check_query_safety(request.query)
if safety.blocked:
yield {
"event": "error",
"data": json.dumps({"message": safety.message}),
}
return
# Emit retrieval start
yield {
"event": "retrieval",
"data": json.dumps({"status": "searching", "query_id": str(query_id)}),
}
# Retrieve context
try:
retrieval_result = await retriever.retrieve(
query=request.query,
language=request.language,
filters=request.filters,
tier=auth.tier,
)
except Exception as exc:
logger.error("Retrieval failed during stream: %s", exc)
yield {
"event": "error",
"data": json.dumps({"message": "Retrieval failed. Please try again."}),
}
return
yield {
"event": "retrieval",
"data": json.dumps({
"status": "found",
"sources": len(retrieval_result.context_blocks),
}),
}
# Emit individual citations
for block in retrieval_result.context_blocks:
yield {
"event": "citation",
"data": json.dumps({
"document": block.get("document_title", ""),
"page": block.get("page_range_start"),
"section": block.get("section_heading", ""),
}),
}
# Build prompt
prompt_messages = build_rag_prompt(
query=request.query,
context_chunks=retrieval_result.context_blocks,
language=request.language,
mode="concise",
)
# Stream LLM response
full_answer = ""
model_used = ""
tokens_used = 0
try:
async for chunk in llm.stream(messages=prompt_messages):
full_answer += chunk.text
model_used = chunk.model_used
tokens_used = chunk.tokens_used
yield {
"event": "delta",
"data": json.dumps({"content": chunk.text}),
}
except Exception as exc:
logger.error("LLM streaming failed: %s", exc)
yield {
"event": "error",
"data": json.dumps({"message": "Generation failed. Please try again."}),
}
return
# Format final citations
formatted = citation_fmt.format_response(
raw_answer=full_answer,
retrieved_chunks=retrieval_result.context_blocks,
)
latency_ms = int((time.monotonic() - start_time) * 1000)
# Log to audit trail
try:
query_log = QueryLog(
id=query_id,
query_text=request.query,
language=request.language,
mode="concise",
user_tier=auth.tier,
user_id=auth.user.id if auth.user else None,
api_key_id=auth.api_key.id if auth.api_key else None,
fingerprint=auth.fingerprint,
filters=request.filters.model_dump() if request.filters else None,
answer=full_answer,
citations=[c.model_dump() for c in formatted.citations] if formatted.citations else None,
chunks_considered=retrieval_result.total_candidates,
latency_ms=latency_ms,
confidence="high" if retrieval_result.top_score > 0.8 else "medium",
llm_model_used=model_used,
llm_tokens_used=tokens_used,
)
db.add(query_log)
await db.commit()
except Exception as exc:
logger.error("Failed to log stream query: %s", exc)
# Emit done event
yield {
"event": "done",
"data": json.dumps({
"query_id": str(query_id),
"citations": [c.model_dump() for c in formatted.citations],
"confidence": "high" if retrieval_result.top_score > 0.8 else "medium",
"latency_ms": latency_ms,
}),
}
return EventSourceResponse(event_generator())