feat(rag): stream grounded RAG answers over SSE with source citations

Phase 03 (Story: Chat RAG Answer — happy path):
- app/rag/retriever.py: top-k cosine search + parent-doc selection with
  per-doc dedupe and BOR_MAX_CONTEXT_CHARS cap ([…truncated…] marker)
- app/rag/prompts.py: locked persona + HIGH/DEFLECT prompt builders
- app/rag/llm.py: LLMError + chat_stream (turbo, temp 0.4, max 700, stream)
- app/api/chat.py: POST /api/chat SSE — delta* then done{deflected,
  sources, suggestions}; query_log row + PLAN §9 per-turn log line;
  structured error event on mid-stream failure, JSON 503 when DB down
- frontend: SSE reader, live bubble streaming, source chips -> /sources.html,
  red role=alert banner, Send button state that always recovers
- fix(scaffold): [hidden] { display: none !important } — .kb-banner's
  display:flex was overriding the hidden attribute (banner always visible)
- tests: unit (retriever/prompts/sse/llm) + integration (real Postgres RAG
  turn, query_log, error + 503 paths, mid-turn failures) + Playwright story
  suite (grounded answer, log row, raw SSE shape); smoke placeholder test
  replaced with the real never-stale-button contract
This commit is contained in:
2026-08-21 17:17:02 -04:00
parent 99c48cbe06
commit 396e4d47fb
14 changed files with 1266 additions and 59 deletions
+152 -20
View File
@@ -1,29 +1,161 @@
"""POST /api/chat — placeholder (phase 01).
"""POST /api/chat — a RAG chat turn streamed over SSE (PLAN §3/§4).
Phase 03 replaces this with the real RAG pipeline and SSE streaming
(PLAN §4 contract: ``delta`` events + final ``done``). The placeholder
keeps the same JSON shape the frontend already consumes, so the UI round-trip
is exercised end-to-end from day one.
Flow (LOCKED A7/A15): embed the question → pgvector cosine top-K chunks →
distinct parent documents (full text, capped) → 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.
The honesty gate (LOW relevance → deflection) lands in phase 04; every
turn in this phase is grounded (``deflected=false``).
"""
from __future__ import annotations
from fastapi import APIRouter
import json
import logging
import time
from collections.abc import AsyncIterator
from typing import Any
from app.schemas import ChatRequest
from fastapi import APIRouter, Depends
from fastapi.responses import JSONResponse, StreamingResponse
from sqlalchemy.orm import Session
router = APIRouter()
from app.config import get_settings
from app.db import db_available, get_db
from app.models import QueryLog
from app.rag.llm import EmbeddingError, LLMClient, LLMError
from app.rag.prompts import build_high_prompt
from app.rag.retriever import retrieve, select_documents
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: <json>\\n\\n`` (PLAN §4)."""
return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n"
@router.post("/chat")
async def chat(request: ChatRequest) -> dict[str, object]:
"""Placeholder answer — no LLM, no DB."""
return {
"ok": True,
"answer": (
"Hey! My neurons are still wiring up — the real Brain "
"(RAG over your docs, powered by aipi) lands in the next "
"phases. Try me again soon! 🧠"
),
"deflected": False,
"sources": [],
}
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 → top-N full parent documents.
try:
chunks = retrieve(db, question_vec)
docs = select_documents(chunks)
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
top_score = chunks[0].score if chunks else 0.0
source_paths = [f"{d.source}/{d.path}" for d in docs]
messages = [
{"role": "system", "content": build_high_prompt(docs)},
{"role": "user", "content": request.message},
]
# 3. Stream the grounded answer.
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=top_score,
chunk_hits=len(chunks),
deflected=False,
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,
top_score,
get_settings().relevance_threshold,
False,
source_paths,
total_ms,
)
yield sse_event(
ChatDoneEvent(
deflected=False,
sources=[
SourceRef(source=d.source, path=d.path, title=d.title) for d in docs
],
suggestions=[],
).model_dump()
)
return StreamingResponse(stream(), media_type="text/event-stream", headers=SSE_HEADERS)