""" 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())