File size: 6,356 Bytes
42a0d15 37b5223 42a0d15 37b5223 42a0d15 37b5223 42a0d15 37b5223 42a0d15 37b5223 42a0d15 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 | """
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())
|