Files
2026-07-31 12:27:41 +08:00

50 lines
1.4 KiB
Python

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()}