Files
GraphRAGAgent/backend/routers/query.py
T
admin ebf27a6c3e fix: KGEmptyError 自定义异常 + CORS 环境变量配置
- qa_service 定义 KGEmptyError 异常,run_query 抛出;routers/query 精确捕获 → HTTP 400 + code 3002,移除字符串匹配
- main.py CORS 改为读 CORS_ORIGINS 环境变量:未设置时 wildcard + credentials=False(合规默认),设置后显式列表 + credentials=True
- .env.example 新增 CORS_ORIGINS 默认值(后端 + 前端 4 个域名)
2026-06-18 14:06:38 +08:00

68 lines
2.0 KiB
Python

"""D 组:QA 问答(4 个端点)"""
import asyncio
from functools import partial
from fastapi import APIRouter
from fastapi.responses import JSONResponse
from models.schemas import APIResponse, BatchQueryRequest, QueryRequest
from services import qa_service as svc
from services.qa_service import KGEmptyError
router = APIRouter(prefix="/query", tags=["QA"])
@router.post("")
async def run_query(body: QueryRequest):
try:
loop = asyncio.get_event_loop()
result = await loop.run_in_executor(
None,
partial(svc.run_query, body.question, [m.model_dump() for m in body.history]),
)
return APIResponse.ok(result)
except KGEmptyError as e:
return JSONResponse(
status_code=400,
content=APIResponse.err(3002, str(e)).model_dump(),
)
except ValueError as e:
return JSONResponse(
status_code=500,
content=APIResponse.err(4001, str(e)).model_dump(),
)
except Exception as e:
return JSONResponse(
status_code=500,
content=APIResponse.err(4001, f"QA service error: {e}").model_dump(),
)
@router.post("/batch", status_code=202)
async def start_batch(body: BatchQueryRequest):
if len(body.questions) > 20:
return JSONResponse(
status_code=400,
content=APIResponse.err(1001, "Maximum 20 questions per batch").model_dump(),
)
result = svc.start_batch(body.questions)
return APIResponse.ok(result)
@router.get("/batch/{batch_id}")
async def get_batch_result(batch_id: str):
result = svc.get_batch_result(batch_id)
if not result:
return JSONResponse(
status_code=404,
content=APIResponse.err(2002, f"Batch '{batch_id}' not found").model_dump(),
)
return APIResponse.ok(result)
@router.get("/history")
async def get_query_history(page: int = 1, page_size: int = 20):
page_size = min(page_size, 50)
result = svc.get_history(page, page_size)
return APIResponse.ok(result)