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:
+152
-20
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user