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