26-7-31-1
This commit is contained in:
@@ -0,0 +1,49 @@
|
||||
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()}
|
||||
Reference in New Issue
Block a user