Files
more_dots/api/endpoints.py
T
2026-02-26 19:23:54 +08:00

161 lines
5.3 KiB
Python

from fastapi import APIRouter, HTTPException, Depends
from fastapi.responses import StreamingResponse
from schemas.agent_input import AgentInput
from schemas.agent_output import AgentOutput
from schemas.tool_input import ToolInput
from schemas.tool_output import ToolOutput
from workflows.workflow_manager import WorkflowType
from api.dependencies import get_workflow_manager, get_nacos_manager, get_service_config, get_tool_router, get_prompt_manager
from services.ragflow_sync import RagflowSync
router = APIRouter()
def _resolve_workflow_type(value: str) -> WorkflowType:
try:
return WorkflowType(value)
except Exception as e:
raise ValueError(f"不支持的工作流类型: {value}") from e
@router.get("/health")
def health_check(service_config=Depends(get_service_config)):
return {
"status": "ok",
"service_name": service_config.service_name,
"model_section": service_config.metadata.get("model_section", "")
}
@router.get("/nacos/status")
def nacos_status(nacos_manager=Depends(get_nacos_manager)):
return nacos_manager.status()
@router.post("/api/workflows", response_model=AgentOutput)
def run_workflow(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)):
try:
workflow_type = _resolve_workflow_type(payload.workflow_type)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
result = workflow_manager.execute_workflow(
workflow_type=workflow_type,
user_input=payload.input,
session_id=payload.session_id,
)
return AgentOutput(
session_id=result["session_id"],
workflow_type=result["workflow_type"],
result=result["result"],
)
@router.post("/api/workflows/stream")
def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)):
try:
workflow_type = _resolve_workflow_type(payload.workflow_type)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
if workflow_type != WorkflowType.CONVERSATION:
raise HTTPException(status_code=400, detail="仅支持对话工作流的流式输出")
def _extract_output_text(result: dict) -> str:
context = (result.get("context") or {}) if isinstance(result, dict) else {}
if "sr_api_result" in context:
return str(context.get("sr_api_result") or "")
messages = result.get("messages") if isinstance(result, dict) else None
if messages:
last = messages[-1]
if hasattr(last, "content"):
return str(last.content or "")
return ""
def event_stream():
try:
result = workflow_manager.execute_workflow(
workflow_type=workflow_type,
user_input=payload.input,
session_id=payload.session_id,
)
text = _extract_output_text(result.get("result") or {})
if not text:
yield "event: end\ndata: [DONE]\n\n"
return
chunk_size = 512
for i in range(0, len(text), chunk_size):
chunk = text[i : i + chunk_size]
yield f"data: {chunk}\n\n"
yield "event: end\ndata: [DONE]\n\n"
except Exception as e:
yield f"event: error\ndata: {str(e)}\n\n"
return StreamingResponse(event_stream(), media_type="text/event-stream")
@router.post("/api/tools/execute", response_model=ToolOutput)
def run_tool(payload: ToolInput, tool_router=Depends(get_tool_router)):
result = tool_router.call(payload.tool_name, payload.payload)
return ToolOutput(**result)
@router.post("/api/prompts/reload")
def reload_prompts(prompt_manager=Depends(get_prompt_manager)):
prompt_manager.reload()
return {"ok": True}
@router.post("/api/ragflow/table-retrieval/reload")
def reload_table_retrieval():
syncer = RagflowSync()
result = syncer.upload_table_retrieval()
return {"ok": True, "result": result}
@router.post("/api/ragflow/table-retrieval/upload")
def upload_table_retrieval():
"""上传表名检索模板文档"""
syncer = RagflowSync()
try:
result = syncer.upload_table_retrieval()
return {"ok": True, "result": result}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@router.put("/api/ragflow/table-retrieval/update")
def update_table_retrieval(config: dict):
"""更新表名检索知识库配置"""
syncer = RagflowSync()
try:
result = syncer.update_dataset(syncer._table_retrieval_dataset_id, config)
return {"ok": True, "result": result}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@router.post("/api/ragflow/sql-gen/upload")
def upload_sql_gen():
"""上传 SQL 生成提示词文档"""
syncer = RagflowSync()
try:
result = syncer.upload_sql_gen()
return {"ok": True, "result": result}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@router.put("/api/ragflow/sql-gen/update")
def update_sql_gen(config: dict):
"""更新 SQL 生成知识库配置"""
syncer = RagflowSync()
try:
result = syncer.update_dataset(syncer._sql_gen_dataset_id, config)
return {"ok": True, "result": result}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))