This commit is contained in:
2026-03-11 23:40:39 +08:00
parent db25d61026
commit e062368ef2
15 changed files with 1592 additions and 229 deletions
+181
View File
@@ -13,6 +13,7 @@ from schemas.tool_input import ToolInput
from schemas.tool_output import ToolOutput
from schemas.chat_message_response import ChatMessageResponseDTO
from schemas.chat_message_request import ChatMessageRequestDTO
from schemas.super_agent import SuperAgentRequest, SuperAgentResponse, SuperAgentStreamEvent
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.app_errors import AppError, ErrorCode
@@ -190,6 +191,54 @@ def run_workflow_stream(payload: ChatMessageRequestDTO, workflow_manager=Depends
return StreamingResponse(event_stream(), media_type="text/event-stream")
@router.get("/api/workflows/list")
def list_workflows(workflow_manager=Depends(get_workflow_manager)):
"""列出所有可用工作流"""
workflows = workflow_manager.get_available_workflows()
result = []
for name in workflows:
info = workflow_manager.get_workflow_info(name)
if info:
result.append(info)
return {"workflows": result}
@router.get("/api/workflows/{workflow_name}")
def get_workflow_detail(workflow_name: str, workflow_manager=Depends(get_workflow_manager)):
"""获取工作流详情"""
info = workflow_manager.get_workflow_info(workflow_name)
if not info:
raise HTTPException(status_code=404, detail=f"工作流不存在: {workflow_name}")
return info
@router.get("/api/tools/list")
def list_tools(tool_router=Depends(get_tool_router)):
"""列出所有可用工具"""
tools = tool_router.list_tools()
result = []
for name in tools:
info = tool_router.get_tool_info(name)
if info:
result.append(info)
return {"tools": result}
@router.get("/api/tools/{tool_name}")
def get_tool_detail(tool_name: str, tool_router=Depends(get_tool_router)):
"""获取工具详情"""
info = tool_router.get_tool_info(tool_name)
if not info:
raise HTTPException(status_code=404, detail=f"工具不存在: {tool_name}")
return info
@router.get("/api/tools/stats")
def get_tools_stats(tool_router=Depends(get_tool_router)):
"""获取工具执行统计"""
return tool_router.get_all_stats()
@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)
@@ -251,3 +300,135 @@ def update_sql_gen():
return {"ok": True, "result": result}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@router.post("/api/super-agent/query", response_model=SuperAgentResponse)
def super_agent_query(payload: SuperAgentRequest, workflow_manager=Depends(get_workflow_manager)):
"""Super Agent 同步查询接口"""
trace_id = uuid.uuid4().hex
slog = get_structured_logger()
slog.log("INFO", "super_agent.query.start", trace_id, {
"query": payload.query[:100],
"workflow_type": payload.workflow_type,
"user_id": payload.user_id,
})
conversation_id = payload.conversation_id or uuid.uuid4().hex
try:
workflow_type = _resolve_workflow_type(payload.workflow_type)
except Exception as e:
slog.log("ERROR", "super_agent.query.invalid_type", trace_id, error_code=ErrorCode.INVALID_WORKFLOW_TYPE.value)
return SuperAgentResponse(
conversation_id=conversation_id,
workflow_type=payload.workflow_type,
status="error",
error=f"不支持的工作流类型: {payload.workflow_type}",
)
try:
result = workflow_manager.execute_workflow(
workflow_type=workflow_type,
user_input=payload.query,
session_id=conversation_id,
)
context = (result.get("result") or {}).get("context") or {}
sql_text = context.get("final_sql")
sr_api_result = context.get("sr_api_result")
slog.log("INFO", "super_agent.query.success", trace_id, {
"conversation_id": conversation_id,
"has_sql": bool(sql_text),
"has_result": bool(sr_api_result),
})
return SuperAgentResponse(
conversation_id=conversation_id,
workflow_type=workflow_type.value,
status="success",
sql=sql_text,
result=str(sr_api_result) if sr_api_result else None,
metadata={"trace_id": trace_id},
)
except Exception as e:
slog.log("ERROR", "super_agent.query.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)})
return SuperAgentResponse(
conversation_id=conversation_id,
workflow_type=payload.workflow_type,
status="error",
error=str(e),
metadata={"trace_id": trace_id},
)
@router.post("/api/super-agent/stream")
def super_agent_stream(payload: SuperAgentRequest, workflow_manager=Depends(get_workflow_manager)):
"""Super Agent 流式查询接口"""
trace_id = uuid.uuid4().hex
slog = get_structured_logger()
stream_cfg = Config.get_section("stream")
progress_interval = float(stream_cfg.get("progress_interval", 0.3))
conversation_id = payload.conversation_id or uuid.uuid4().hex
def _build_sse_event(event: str, data: str) -> str:
dto = SuperAgentStreamEvent(
conversation_id=conversation_id,
event=event,
data=data,
timestamp=int(time.time() * 1000),
)
return f"event: {event}\ndata: {json.dumps(dto.model_dump(), ensure_ascii=False)}\n\n"
async def event_stream():
try:
slog.log("INFO", "super_agent.stream.start", trace_id, {
"query": payload.query[:100],
"user_id": payload.user_id,
})
workflow_type = _resolve_workflow_type(payload.workflow_type)
result = await asyncio.to_thread(
workflow_manager.execute_workflow,
workflow_type,
payload.query,
conversation_id,
skip_sr_api=True,
)
context = (result.get("result") or {}).get("context") or {}
sql_text = context.get("final_sql")
if not sql_text:
slog.log("ERROR", "super_agent.stream.sql_failed", trace_id, error_code=ErrorCode.SQL_GENERATION_FAILED.value)
yield _build_sse_event("error", "SQL 生成失败")
yield _build_sse_event("done", "")
return
yield _build_sse_event("sql_generated", sql_text)
yield _build_sse_event("sql_executing", "")
tool = SrApiQueryTool()
task = asyncio.create_task(
asyncio.to_thread(tool.run, json.dumps({"sql": sql_text}, ensure_ascii=False))
)
while not task.done():
yield _build_sse_event("sql_executing", "")
await asyncio.sleep(progress_interval)
sql_result = await task
slog.log("INFO", "super_agent.stream.success", trace_id, {"result_len": len(str(sql_result))})
yield _build_sse_event("result", str(sql_result))
except Exception as e:
slog.log("ERROR", "super_agent.stream.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)})
yield _build_sse_event("error", str(e))
yield _build_sse_event("done", "")
return StreamingResponse(event_stream(), media_type="text/event-stream")