"""POST /api/chat — a RAG chat turn streamed over SSE (PLAN §3/§4). Flow (LOCKED A7/A15): embed the question → pgvector cosine top-K chunks → the **honesty gate** (A8: best score < ``BOR_RELEVANCE_THRESHOLD`` ⇒ deflection) → locked persona prompt (PLAN §6) → ``turbo`` streamed as ``delta`` events → final ``done`` event (``deflected``, ``sources``, ``suggestions``) + ``query_log`` row + the per-turn log line (PLAN §9). Mid-stream failures become a structured ``error`` event; a pre-stream DB outage is a plain 503 JSON. Honesty gate: a weak retrieval (score strictly below the threshold — or an empty KB) flips the turn to deflection mode: the LOW prompt carries weak-hit *titles only* (never document content) plus deterministic "Maybe try" chips, and the ``done`` event / ``query_log`` row record ``deflected=true`` with the weak score. """ from __future__ import annotations import json import logging import time from collections.abc import AsyncIterator, Sequence from dataclasses import dataclass from typing import Any from fastapi import APIRouter, Depends from fastapi.responses import JSONResponse, StreamingResponse from sqlalchemy.orm import Session from app.config import Settings, get_settings from app.db import db_available, get_db from app.models import Document, QueryLog from app.rag.llm import EmbeddingError, LLMClient, LLMError from app.rag.prompts import build_deflect_prompt, build_high_prompt from app.rag.retriever import RetrievedChunk, retrieve, select_documents, weak_hit_titles from app.rag.suggestions import derive_suggestions from app.schemas import ChatDoneEvent, ChatRequest, SourceRef logger = logging.getLogger("app.chat") router = APIRouter(tags=["chat"]) #: Streaming hints: no proxy buffering, no client caching (PLAN A15). SSE_HEADERS = {"Cache-Control": "no-cache", "X-Accel-Buffering": "no"} _llm: LLMClient | None = None def get_llm() -> LLMClient: """Shared LLM client (FastAPI dependency so tests can override it).""" global _llm if _llm is None: _llm = LLMClient(get_settings()) return _llm def sse_event(payload: dict[str, Any]) -> str: """Serialize one SSE frame: ``data: \\n\\n`` (PLAN §4).""" return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n" @dataclass class TurnPlan: """What one chat turn sends to the LLM and reports on ``done``.""" top_score: float deflected: bool system_prompt: str docs: list[Document] # cited sources (weak hits when deflected) suggestions: list[str] # "Maybe try" chips (deflected turns only) def plan_turn(chunks: Sequence[RetrievedChunk], settings: Settings) -> TurnPlan: """Apply the honesty gate (A8) and assemble prompt + context for a turn. * ``top_score >= threshold`` → grounded: HIGH prompt with the full top-N documents, no suggestions. A score exactly at the threshold is an answer — the gate is strict (``score < threshold``). * ``top_score < threshold`` (or no hits at all) → deflected: LOW prompt (``DEFLECT_MODE``) with weak-hit titles only — never document content — plus deterministic alternative-question chips derived from those titles. """ top_score = chunks[0].score if chunks else 0.0 if top_score >= settings.relevance_threshold: docs = select_documents( chunks, n=settings.top_n_docs, max_chars=settings.max_context_chars ) return TurnPlan(top_score, False, build_high_prompt(docs), docs, []) titles = weak_hit_titles(chunks) return TurnPlan( top_score, True, build_deflect_prompt(titles), select_documents(chunks, n=settings.top_n_docs, max_chars=settings.max_context_chars), derive_suggestions(titles, settings.suggestions), ) @router.post("/chat") async def chat( request: ChatRequest, db: Session = Depends(get_db), # noqa: B008 llm: LLMClient = Depends(get_llm), # noqa: B008 ): """One chat turn: SSE stream of ``delta`` events + a final ``done``.""" if not db_available(): return JSONResponse( status_code=503, content={ "detail": ( "The knowledge base is offline — start Postgres with " "`podman compose up -d db`, then ask again." ) }, ) started = time.monotonic() async def stream() -> AsyncIterator[str]: # 1. Embed the question. t0 = time.monotonic() try: question_vec = await llm.embed_one(request.message) except EmbeddingError as e: logger.error("chat: embedding failed question=%r — %s", request.message, e) yield sse_event( { "type": "error", "detail": "I couldn't reach the embedding model — please try again.", } ) return embed_ms = int((time.monotonic() - t0) * 1000) # 2. Retrieve top-K chunks, then the honesty gate (A8) picks the # HIGH (grounded) or LOW (deflected) prompt + context. settings = get_settings() try: chunks = retrieve(db, question_vec) plan = plan_turn(chunks, settings) except Exception: # noqa: BLE001 — DB failure mid-turn logger.exception("chat: retrieval failed question=%r", request.message) yield sse_event( { "type": "error", "detail": "The knowledge base went offline mid-question — is Postgres up?", } ) return source_paths = [f"{d.source}/{d.path}" for d in plan.docs] messages = [ {"role": "system", "content": plan.system_prompt}, {"role": "user", "content": request.message}, ] # 3. Stream the answer (grounded, or an honest deflection). try: async for piece in llm.chat_stream(messages): yield sse_event({"type": "delta", "text": piece}) except LLMError as e: logger.error("chat: LLM stream failed question=%r — %s", request.message, e) yield sse_event( {"type": "error", "detail": "The chat model dropped the connection — try again?"} ) return # 4. Durable record + required per-turn log line (PLAN §9). total_ms = int((time.monotonic() - started) * 1000) try: db.add( QueryLog( question=request.message, top_score=plan.top_score, chunk_hits=len(chunks), deflected=plan.deflected, sources=", ".join(source_paths), latency_ms=total_ms, ) ) db.commit() except Exception: # noqa: BLE001 — the answer already went out logger.exception("chat: failed to write query_log question=%r", request.message) logger.info( "question=%r embed_ms=%d top_score=%.3f threshold=%.2f deflected=%s " "sources=%r total_ms=%d", request.message, embed_ms, plan.top_score, settings.relevance_threshold, plan.deflected, source_paths, total_ms, ) yield sse_event( ChatDoneEvent( deflected=plan.deflected, sources=[ SourceRef(source=d.source, path=d.path, title=d.title) for d in plan.docs ], suggestions=plan.suggestions, ).model_dump() ) return StreamingResponse(stream(), media_type="text/event-stream", headers=SSE_HEADERS)