147 lines
4.9 KiB
Python
147 lines
4.9 KiB
Python
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,
|
|
)
|