from __future__ import annotations from typing import Annotated, Any, Protocol from fastapi import APIRouter, Depends, HTTPException, Request from pydantic import BaseModel, Field from rag.answer import GroundedAnswerService from rag.models import QueryIntent, SubjectScope class TraceWriter(Protocol): def save(self, **fields: Any) -> str: ... class RagQueryRequest(BaseModel): query: str = Field(min_length=1, max_length=4000) subject_scope: SubjectScope intent: QueryIntent # Optional: when present, the turn is answered in conversation context # (follow-up inheritance, clarify, smalltalk). Absent → single-turn, exactly # as before, so existing callers are unchanged. conversation_id: str | None = Field(default=None, max_length=128) class CitationResponse(BaseModel): chunk_id: str printed_page_start: int printed_page_end: int physical_page: int block_id: str | None = None bbox: tuple[float, float, float, float] | None = None source_crop: str | None = None attachment: str | None = None class RagQueryResponse(BaseModel): trace_id: str decision: str reason: str answer: str | None resolved_drug_id: str | None citations: list[CitationResponse] def _answer_service(request: Request) -> GroundedAnswerService: service = getattr(request.app.state, "answer_service", None) if service is None: raise HTTPException(status_code=503, detail="RAG backend is not configured") return service def _trace_writer(request: Request) -> TraceWriter: writer = getattr(request.app.state, "trace_writer", None) if writer is None: raise HTTPException(status_code=503, detail="Trace database is not configured") return writer router = APIRouter(prefix="/v1/rag", tags=["rag"]) class SuggestResponse(BaseModel): suggestions: list[str] @router.get("/suggest", response_model=SuggestResponse) def suggest_drugs(q: str, request: Request) -> SuggestResponse: """As-you-type drug-name autocomplete, so a name is picked, not mistyped.""" conversational = getattr(request.app.state, "conversational", None) if conversational is None or not q.strip(): return SuggestResponse(suggestions=[]) return SuggestResponse(suggestions=conversational.complete(q.strip())) def _map_citations(items) -> list[CitationResponse]: return [ CitationResponse( chunk_id=item.chunk_id, printed_page_start=item.printed_page_start, printed_page_end=item.printed_page_end, physical_page=item.physical_page, block_id=item.block_id, bbox=item.bbox, source_crop=item.source_crop, attachment=item.attachment, ) for item in items ] @router.post("/query", response_model=RagQueryResponse) def query_rag( payload: RagQueryRequest, request: Request, answers: Annotated[GroundedAnswerService, Depends(_answer_service)], traces: Annotated[TraceWriter, Depends(_trace_writer)], ) -> RagQueryResponse: conversational = getattr(request.app.state, "conversational", None) # Single-turn path (no conversation id, or conversational layer disabled): # unchanged behaviour so existing callers keep working. if payload.conversation_id is None or conversational is None: grounded = answers.answer(payload.query, payload.subject_scope, payload.intent) decision = grounded.result.decision.value reason = grounded.result.reason answer = grounded.answer resolved_drug_id = grounded.result.resolved_drug_id citations = _map_citations(grounded.citations) else: turn = conversational.answer( payload.conversation_id, payload.query, payload.subject_scope, payload.intent, ) if turn.clarification is not None: decision, reason = "clarify", turn.clarification.reason answer, resolved_drug_id, citations = turn.clarification.question, None, [] elif turn.grounded is not None: decision = turn.grounded.result.decision.value reason = turn.grounded.result.reason answer = turn.answer resolved_drug_id = turn.grounded.result.resolved_drug_id citations = _map_citations(turn.grounded.citations) else: # smalltalk decision, reason = "answerable", turn.reason answer, resolved_drug_id, citations = turn.answer, None, [] trace_id = traces.save( query=payload.query, subject_scope=payload.subject_scope.value, intent=payload.intent.value, decision=decision, reason=reason, resolved_drug_id=resolved_drug_id, citations=tuple(item.model_dump() for item in citations), ) return RagQueryResponse( trace_id=trace_id, decision=decision, reason=reason, answer=answer, resolved_drug_id=resolved_drug_id, citations=citations, )