from fastapi import APIRouter from pydantic import BaseModel, Field from app.schemas.models import RagQueryRequest, RagQueryResponse, SourceRef from app.services.rag_pipeline import get_rag router = APIRouter(prefix="/rag", tags=["rag"]) class IngestRequest(BaseModel): title: str content: str category: str = "自定义" class IngestResponse(BaseModel): ok: bool = True chunks: int = 0 @router.post("/query", response_model=RagQueryResponse) def rag_query(body: RagQueryRequest) -> RagQueryResponse: rag = get_rag() result = rag.query(body.query, top_k=body.top_k, extra_context=body.context or "") raw_sources = result.get("sources") or [] sources: list[SourceRef] = [] for s in raw_sources: if isinstance(s, SourceRef): sources.append(s) elif isinstance(s, dict): sources.append(SourceRef(**s)) return RagQueryResponse( answer=result.get("answer") or "", sources=sources, engine=result.get("engine") or "langchain-rag", ) @router.post("/ingest", response_model=IngestResponse) def rag_ingest(body: IngestRequest) -> IngestResponse: rag = get_rag() rag.ingest_text(body.title, body.content, body.category) return IngestResponse(ok=True, chunks=rag.doc_count()) @router.get("/stats") def rag_stats() -> dict: rag = get_rag() return {"chunks": rag.doc_count()}