This commit is contained in:
2026-03-02 17:41:04 +08:00
parent 460c2e87b8
commit c9bebe5615
7 changed files with 322 additions and 86 deletions
+14 -17
View File
@@ -12,6 +12,7 @@ from schemas.agent_output import AgentOutput
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 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
@@ -68,8 +69,8 @@ def run_workflow(payload: AgentInput, workflow_manager=Depends(get_workflow_mana
try:
result = workflow_manager.execute_workflow(
workflow_type=workflow_type,
user_input=payload.input,
session_id=payload.session_id,
user_input=payload.query,
session_id=payload.conversation_id,
)
slog.log("INFO", "run_workflow.success", trace_id, {"session_id": result.get("session_id")})
except Exception as e:
@@ -96,8 +97,8 @@ def generate_sql(payload: AgentInput, workflow_manager=Depends(get_workflow_mana
result = workflow_manager.execute_workflow(
workflow_type=workflow_type,
user_input=payload.input,
session_id=payload.session_id,
user_input=payload.query,
session_id=payload.conversation_id,
skip_sr_api=True,
)
@@ -118,16 +119,12 @@ def generate_sql(payload: AgentInput, workflow_manager=Depends(get_workflow_mana
@router.post("/api/workflows/stream")
def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)):
def run_workflow_stream(payload: ChatMessageRequestDTO, workflow_manager=Depends(get_workflow_manager)):
trace_id = uuid.uuid4().hex
slog = get_structured_logger()
try:
workflow_type = _resolve_workflow_type(payload.workflow_type)
except Exception as e:
raise _to_http_error(e)
if workflow_type != WorkflowType.CONVERSATION:
raise _to_http_error(AppError(code=ErrorCode.INVALID_WORKFLOW_TYPE, message="仅支持对话工作流的流式输出", status_code=400))
if payload.response_mode != "streaming":
raise _to_http_error(AppError(code=ErrorCode.INVALID_WORKFLOW_TYPE, message="/api/workflows/stream 仅支持 response_mode=streaming", status_code=400))
stream_cfg = Config.get_section("stream")
progress_interval = float(stream_cfg.get("progress_interval", 0.3))
@@ -147,16 +144,16 @@ def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workfl
async def event_stream():
try:
slog.log("INFO", "stream.start", trace_id, {"workflow_type": payload.workflow_type})
slog.log("INFO", "stream.start", trace_id, {"workflow_type": WorkflowType.CONVERSATION.value})
# 1) 先仅生成 SQL(不执行 SR API)
result = await asyncio.to_thread(
workflow_manager.execute_workflow,
workflow_type,
payload.input,
payload.session_id,
WorkflowType.CONVERSATION,
payload.query,
payload.conversation_id,
skip_sr_api=True,
)
conversation_id = str(result.get("session_id") or payload.session_id or task_id)
conversation_id = str(result.get("session_id") or payload.conversation_id or task_id)
context = (result.get("result") or {}).get("context") or {}
sql_text = str(context.get("final_sql") or "")
@@ -186,7 +183,7 @@ def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workfl
yield "event: end\ndata: [DONE]\n\n"
except Exception as e:
slog.log("ERROR", "stream.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)})
conversation_id = str(payload.session_id or task_id)
conversation_id = str(payload.conversation_id or task_id)
yield _build_message(conversation_id, str(e))
yield "event: end\ndata: [DONE]\n\n"