2783 lines
113 KiB
Python
2783 lines
113 KiB
Python
"""
|
||
AI 工作流服务
|
||
|
||
数据权限:
|
||
- 使用 list_workflows_with_data_scope() 自动应用数据权限
|
||
- 支持本人、本部门、本部门及下级、全部等数据范围
|
||
"""
|
||
import copy
|
||
import logging
|
||
import time
|
||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||
from queue import Queue
|
||
from threading import Lock
|
||
from datetime import datetime
|
||
from typing import Any, Dict, Generator, List, Optional
|
||
|
||
from sqlalchemy import select, func, and_
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from ai_platform.models import AIWorkflow, AIWorkflowVersion, AIWorkflowRun
|
||
from ai_platform.nodes import NodeContext, NodeRegistry, NodeResult
|
||
from app.data_scope_utils import get_data_scope_filter, apply_data_scope_to_conditions
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
def _apply_node_result_warnings(event: Dict[str, Any], result: NodeResult) -> Dict[str, Any]:
|
||
"""将节点 metadata.warnings 透传到 SSE node_complete 事件"""
|
||
warnings = (result.metadata or {}).get('warnings')
|
||
if warnings:
|
||
event['warnings'] = warnings
|
||
return event
|
||
|
||
|
||
def _resolve_workflow_definition(workflow: AIWorkflow, use_draft: bool) -> dict:
|
||
if use_draft:
|
||
return workflow.definition or {}
|
||
return workflow.published_definition or workflow.definition or {}
|
||
|
||
|
||
def _resolve_trigger_type(use_draft: bool, trigger_type: Optional[str] = None) -> str:
|
||
if trigger_type:
|
||
return trigger_type
|
||
return 'editor_draft' if use_draft else 'editor_published'
|
||
|
||
|
||
def _node_label_from_map(node_map: Dict[str, Any], node_id: str) -> str:
|
||
node = node_map.get(node_id) or {}
|
||
data = node.get('data') or {}
|
||
return data.get('label') or node.get('label') or node_id
|
||
|
||
|
||
def _safe_preview(value: Any, max_length: int = 220) -> str:
|
||
if value is None or value == '':
|
||
return ''
|
||
if isinstance(value, str):
|
||
text = value
|
||
else:
|
||
import json
|
||
|
||
try:
|
||
text = json.dumps(value, ensure_ascii=False)
|
||
except TypeError:
|
||
text = str(value)
|
||
text = text.strip()
|
||
return text if len(text) <= max_length else f'{text[:max_length]}...'
|
||
|
||
|
||
def _node_collaboration_mode(node_type: str, data: Dict[str, Any]) -> str:
|
||
if data.get('agent_code'):
|
||
return 'agent'
|
||
if node_type == 'parallel':
|
||
return 'parallel'
|
||
if node_type == 'subflow':
|
||
return 'subflow'
|
||
if node_type == 'intent':
|
||
return 'router'
|
||
if node_type in ('confirm', 'question', 'choice'):
|
||
return 'human_input'
|
||
if node_type == 'llm':
|
||
return 'llm'
|
||
return 'workflow'
|
||
|
||
|
||
def _node_collaboration_metadata(node_map: Dict[str, Any], node_id: str, node_type: str) -> dict:
|
||
node = node_map.get(node_id) or {}
|
||
data = node.get('data') or {}
|
||
metadata: Dict[str, Any] = {
|
||
'node_id': node_id,
|
||
'node_type': node_type,
|
||
'node_label': _node_label_from_map(node_map, node_id),
|
||
'collaboration_role': data.get('label') or node_id,
|
||
'collaboration_mode': _node_collaboration_mode(node_type, data),
|
||
}
|
||
|
||
agent_code = (data.get('agent_code') or '').strip()
|
||
if agent_code:
|
||
metadata['agent_code'] = agent_code
|
||
if data.get('agent_name'):
|
||
metadata['agent_name'] = data.get('agent_name')
|
||
|
||
model_id = data.get('model_id')
|
||
if model_id:
|
||
metadata['model_id'] = str(model_id)
|
||
|
||
if data.get('subflow_name'):
|
||
metadata['subflow_name'] = data.get('subflow_name')
|
||
if data.get('show_subflow_messages') is not None:
|
||
metadata['show_subflow_messages'] = data.get('show_subflow_messages')
|
||
if data.get('forward_interactive') is not None:
|
||
metadata['forward_interactive'] = data.get('forward_interactive')
|
||
|
||
branches = data.get('branches') or []
|
||
if node_type == 'parallel' and branches:
|
||
metadata['branches'] = [
|
||
{
|
||
'id': branch.get('id'),
|
||
'name': branch.get('name') or branch.get('id'),
|
||
}
|
||
for branch in branches
|
||
if branch.get('id')
|
||
]
|
||
|
||
return metadata
|
||
|
||
|
||
def _build_communication_metadata(
|
||
node_map: Dict[str, Any],
|
||
node_id: str,
|
||
node_type: str,
|
||
*,
|
||
event_type: str,
|
||
status: Optional[str] = None,
|
||
branch_id: Optional[str] = None,
|
||
branch_label: Optional[str] = None,
|
||
target_node_id: Optional[str] = None,
|
||
output: Any = None,
|
||
error: Any = None,
|
||
waiting_config: Optional[Dict[str, Any]] = None,
|
||
) -> dict:
|
||
actor = _node_collaboration_metadata(node_map, node_id, node_type)
|
||
message = _safe_preview(error or output)
|
||
if waiting_config:
|
||
message = (
|
||
waiting_config.get('title')
|
||
or waiting_config.get('question')
|
||
or waiting_config.get('content')
|
||
or message
|
||
)
|
||
|
||
communication: Dict[str, Any] = {
|
||
'event': event_type,
|
||
'status': status or '',
|
||
'actor': {
|
||
key: actor.get(key)
|
||
for key in (
|
||
'node_id',
|
||
'node_type',
|
||
'node_label',
|
||
'agent_code',
|
||
'agent_name',
|
||
'collaboration_role',
|
||
'collaboration_mode',
|
||
)
|
||
if actor.get(key)
|
||
},
|
||
'message': _safe_preview(message),
|
||
}
|
||
if branch_id:
|
||
communication['channel'] = 'parallel_branch'
|
||
communication['branch_id'] = branch_id
|
||
communication['branch_label'] = branch_label or branch_id
|
||
elif node_type == 'subflow':
|
||
communication['channel'] = 'subflow'
|
||
elif node_type in ('confirm', 'question', 'choice') or event_type == 'waiting_input':
|
||
communication['channel'] = 'human_input'
|
||
else:
|
||
communication['channel'] = 'workflow_edge'
|
||
if target_node_id:
|
||
target_type = (node_map.get(target_node_id) or {}).get('type', '')
|
||
target = _node_collaboration_metadata(node_map, target_node_id, target_type)
|
||
communication['target'] = {
|
||
key: target.get(key)
|
||
for key in (
|
||
'node_id',
|
||
'node_type',
|
||
'node_label',
|
||
'agent_code',
|
||
'agent_name',
|
||
'collaboration_role',
|
||
'collaboration_mode',
|
||
)
|
||
if target.get(key)
|
||
}
|
||
return communication
|
||
|
||
|
||
def _resolve_handoff_target(
|
||
current_node_id: str,
|
||
result: Optional[NodeResult],
|
||
edge_map: Dict[str, List[str]],
|
||
parallel_edge_map: Dict[str, Dict[str, str]],
|
||
node_map: Dict[str, Any],
|
||
) -> Optional[str]:
|
||
if result and result.next_node_id:
|
||
if result.next_node_id in parallel_edge_map.get(current_node_id, {}):
|
||
return parallel_edge_map[current_node_id][result.next_node_id]
|
||
if result.next_node_id in node_map:
|
||
return result.next_node_id
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
return next_nodes[0] if next_nodes else None
|
||
|
||
|
||
def _make_execution_log_entry(
|
||
node_map: Dict[str, Any],
|
||
node_id: str,
|
||
node_type: str,
|
||
**extra: Any,
|
||
) -> dict:
|
||
extra.setdefault('timestamp', datetime.now().isoformat())
|
||
metadata = dict(extra.pop('metadata', {}) or {})
|
||
collaboration = _node_collaboration_metadata(node_map, node_id, node_type)
|
||
branch_id = extra.get('branch') or extra.get('branch_id')
|
||
branch_label = extra.get('branch_label')
|
||
if branch_id:
|
||
collaboration['branch_id'] = branch_id
|
||
collaboration['branch_label'] = branch_label or branch_id
|
||
if collaboration:
|
||
metadata['collaboration'] = collaboration
|
||
metadata['communication'] = _build_communication_metadata(
|
||
node_map,
|
||
node_id,
|
||
node_type,
|
||
event_type='node_complete',
|
||
status=extra.get('status'),
|
||
branch_id=branch_id,
|
||
branch_label=branch_label,
|
||
target_node_id=extra.get('target_node_id'),
|
||
output=extra.get('output'),
|
||
error=extra.get('error'),
|
||
)
|
||
metadata['communication']['timestamp'] = extra.get('timestamp')
|
||
return {
|
||
'node_id': node_id,
|
||
'node_type': node_type,
|
||
'node_label': _node_label_from_map(node_map, node_id),
|
||
'metadata': metadata,
|
||
**extra,
|
||
}
|
||
|
||
|
||
def _node_result_metadata(result: NodeResult) -> dict:
|
||
return dict(result.metadata or {})
|
||
|
||
|
||
def _event_collaboration_metadata(
|
||
node_map: Dict[str, Any],
|
||
node_id: str,
|
||
node_type: str,
|
||
result: Optional[NodeResult] = None,
|
||
) -> dict:
|
||
metadata = _node_collaboration_metadata(node_map, node_id, node_type)
|
||
if not result:
|
||
return metadata
|
||
|
||
result_metadata = result.metadata or {}
|
||
for key in (
|
||
'agent_code',
|
||
'agent_name',
|
||
'model',
|
||
'model_id',
|
||
'output_variable',
|
||
'subflow_name',
|
||
'prompt_tokens',
|
||
'completion_tokens',
|
||
'confidence',
|
||
'matched_branch',
|
||
):
|
||
value = result_metadata.get(key)
|
||
if value is not None and value != '':
|
||
metadata[key] = value
|
||
return metadata
|
||
|
||
|
||
def _parallel_branch_labels(node_map: Dict[str, Any], parallel_node_id: str) -> Dict[str, str]:
|
||
node = node_map.get(parallel_node_id) or {}
|
||
branches = (node.get('data') or {}).get('branches') or []
|
||
labels: Dict[str, str] = {}
|
||
for branch in branches:
|
||
branch_id = branch.get('id')
|
||
if branch_id:
|
||
labels[branch_id] = branch.get('name') or branch_id
|
||
return labels
|
||
|
||
|
||
def _merge_event_collaboration(
|
||
event: Dict[str, Any],
|
||
node_map: Dict[str, Any],
|
||
node_id: str,
|
||
node_type: str,
|
||
result: Optional[NodeResult] = None,
|
||
) -> Dict[str, Any]:
|
||
collaboration = _event_collaboration_metadata(node_map, node_id, node_type, result)
|
||
if event.get('branch_id'):
|
||
collaboration['branch_id'] = event.get('branch_id')
|
||
collaboration['branch_label'] = event.get('branch_label') or event.get('branch_id')
|
||
if collaboration:
|
||
event['collaboration'] = collaboration
|
||
event_error = None
|
||
if event.get('type') == 'error':
|
||
event_error = event.get('error_message') or event.get('message') or event.get('content')
|
||
event['communication'] = _build_communication_metadata(
|
||
node_map,
|
||
node_id,
|
||
node_type,
|
||
event_type=event.get('type') or '',
|
||
status=event.get('status'),
|
||
branch_id=event.get('branch_id'),
|
||
branch_label=event.get('branch_label'),
|
||
target_node_id=event.get('target_node_id'),
|
||
output=(event.get('outputs') or {}).get('output') if isinstance(event.get('outputs'), dict) else event.get('outputs'),
|
||
error=event_error,
|
||
waiting_config=event.get('waiting_config') or event.get('config'),
|
||
)
|
||
return event
|
||
|
||
|
||
def _snapshot_node_inputs(context: NodeContext) -> dict:
|
||
return {
|
||
'variables': dict(context.variables),
|
||
'previous_output': context.previous_output,
|
||
}
|
||
|
||
# 资源类型(用于数据权限配置)
|
||
RESOURCE_TYPE = "ai_workflow"
|
||
RESOURCE_DISPLAY_NAME = "AI工作流"
|
||
|
||
|
||
class AIWorkflowService:
|
||
"""
|
||
AI 工作流服务
|
||
|
||
管理工作流的执行
|
||
"""
|
||
|
||
def __init__(self, db: Optional[AsyncSession] = None):
|
||
self._db = db
|
||
|
||
def _create_workflow_run_record(
|
||
self,
|
||
workflow: AIWorkflow,
|
||
inputs: Dict[str, Any],
|
||
*,
|
||
app_id: Optional[str] = None,
|
||
conversation_id: Optional[str] = None,
|
||
use_draft: bool = False,
|
||
trigger_type: Optional[str] = None,
|
||
) -> AIWorkflowRun:
|
||
from utils.context import get_current_user_id_from_context
|
||
|
||
definition = _resolve_workflow_definition(workflow, use_draft)
|
||
return AIWorkflowRun(
|
||
workflow_id=workflow.id,
|
||
app_id=app_id,
|
||
conversation_id=conversation_id,
|
||
user_id=get_current_user_id_from_context(),
|
||
status='running',
|
||
inputs=inputs,
|
||
trigger_type=_resolve_trigger_type(use_draft, trigger_type),
|
||
use_draft=use_draft,
|
||
workflow_version=workflow.published_version if not use_draft else None,
|
||
definition_snapshot=copy.deepcopy(definition),
|
||
started_at=datetime.now(),
|
||
)
|
||
|
||
async def get_workflow(self, workflow_id: str) -> Optional[AIWorkflow]:
|
||
"""获取工作流"""
|
||
if not self._db:
|
||
return None
|
||
result = await self._db.execute(
|
||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def get_workflow_by_code(self, code: str) -> Optional[AIWorkflow]:
|
||
"""根据编码获取工作流"""
|
||
if not self._db:
|
||
return None
|
||
result = await self._db.execute(
|
||
select(AIWorkflow).where(AIWorkflow.code == code, AIWorkflow.is_deleted == False)
|
||
)
|
||
return result.scalar_one_or_none()
|
||
|
||
async def list_workflows(
|
||
self,
|
||
page: int = 1,
|
||
page_size: int = 20,
|
||
status: str = None,
|
||
name: str = None,
|
||
) -> tuple:
|
||
"""
|
||
获取工作流列表
|
||
|
||
Returns:
|
||
(workflows, total)
|
||
"""
|
||
if not self._db:
|
||
return [], 0
|
||
|
||
query = select(AIWorkflow).where(AIWorkflow.is_deleted == False)
|
||
|
||
if status:
|
||
query = query.where(AIWorkflow.status == status)
|
||
if name:
|
||
query = query.where(AIWorkflow.name.ilike(f"%{name}%"))
|
||
|
||
# 获取总数
|
||
count_result = await self._db.execute(
|
||
select(func.count()).select_from(query.subquery())
|
||
)
|
||
total = count_result.scalar() or 0
|
||
|
||
# 分页
|
||
offset = (page - 1) * page_size
|
||
query = query.offset(offset).limit(page_size)
|
||
|
||
result = await self._db.execute(query)
|
||
workflows = list(result.scalars().all())
|
||
|
||
return workflows, total
|
||
|
||
async def list_workflows_with_data_scope(
|
||
self,
|
||
page: int = 1,
|
||
page_size: int = 20,
|
||
status: str = None,
|
||
name: str = None,
|
||
workflow_type: str = None,
|
||
application_id: str = None,
|
||
) -> tuple:
|
||
"""
|
||
获取工作流列表(带数据权限过滤)
|
||
|
||
自动从上下文获取当前用户信息,应用数据权限过滤
|
||
|
||
Returns:
|
||
(workflows, total)
|
||
"""
|
||
if not self._db:
|
||
return [], 0
|
||
|
||
conditions = [AIWorkflow.is_deleted == False]
|
||
|
||
if status:
|
||
conditions.append(AIWorkflow.status == status)
|
||
if name:
|
||
conditions.append(AIWorkflow.name.ilike(f"%{name}%"))
|
||
if workflow_type:
|
||
conditions.append(AIWorkflow.workflow_type == workflow_type)
|
||
if application_id:
|
||
conditions.append(AIWorkflow.application_id == application_id)
|
||
else:
|
||
conditions.append(AIWorkflow.application_id.is_(None))
|
||
|
||
# 获取数据权限过滤条件并应用
|
||
data_scope_filter = await get_data_scope_filter(self._db, RESOURCE_TYPE)
|
||
scope_conditions = apply_data_scope_to_conditions(AIWorkflow, data_scope_filter)
|
||
conditions.extend(scope_conditions)
|
||
|
||
# 获取总数
|
||
query = select(AIWorkflow).where(and_(*conditions))
|
||
count_result = await self._db.execute(
|
||
select(func.count()).select_from(query.subquery())
|
||
)
|
||
total = count_result.scalar() or 0
|
||
|
||
# 分页
|
||
offset = (page - 1) * page_size
|
||
query = query.order_by(AIWorkflow.sort.desc(), AIWorkflow.sys_create_datetime.desc())
|
||
query = query.offset(offset).limit(page_size)
|
||
|
||
result = await self._db.execute(query)
|
||
workflows = list(result.scalars().all())
|
||
|
||
return workflows, total
|
||
|
||
async def create_workflow(self, data: Dict) -> AIWorkflow:
|
||
"""创建工作流"""
|
||
if not self._db:
|
||
raise ValueError("数据库会话未初始化")
|
||
|
||
workflow = AIWorkflow(
|
||
name=data.get('name', ''),
|
||
code=data.get('code', ''),
|
||
description=data.get('description', ''),
|
||
definition=data.get('definition', {}),
|
||
input_variables=data.get('input_variables', []),
|
||
output_variables=data.get('output_variables', []),
|
||
)
|
||
|
||
# 自动填充创建人和部门
|
||
from utils.context import get_current_user_info_from_context
|
||
user_info = get_current_user_info_from_context()
|
||
if user_info:
|
||
if not workflow.sys_creator_id:
|
||
workflow.sys_creator_id = user_info.get('user_id')
|
||
if not workflow.sys_dept_id and user_info.get('dept_id'):
|
||
workflow.sys_dept_id = user_info.get('dept_id')
|
||
|
||
self._db.add(workflow)
|
||
await self._db.commit()
|
||
await self._db.refresh(workflow)
|
||
return workflow
|
||
|
||
async def update_workflow(self, workflow_id: str, data: Dict) -> Optional[AIWorkflow]:
|
||
"""更新工作流"""
|
||
if not self._db:
|
||
return None
|
||
|
||
result = await self._db.execute(
|
||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||
)
|
||
workflow = result.scalar_one_or_none()
|
||
if not workflow:
|
||
return None
|
||
|
||
for key, value in data.items():
|
||
if hasattr(workflow, key) and key not in ('id', 'code'):
|
||
setattr(workflow, key, value)
|
||
|
||
await self._db.commit()
|
||
await self._db.refresh(workflow)
|
||
return workflow
|
||
|
||
async def delete_workflow(self, workflow_id: str) -> bool:
|
||
"""删除工作流"""
|
||
if not self._db:
|
||
return False
|
||
|
||
result = await self._db.execute(
|
||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||
)
|
||
workflow = result.scalar_one_or_none()
|
||
if not workflow:
|
||
return False
|
||
|
||
workflow.is_deleted = True
|
||
await self._db.commit()
|
||
return True
|
||
|
||
async def publish_workflow(self, workflow_id: str, description: str = '') -> Optional[AIWorkflow]:
|
||
"""
|
||
发布工作流
|
||
|
||
1. 创建版本记录
|
||
2. 更新工作流的发布状态
|
||
"""
|
||
if not self._db:
|
||
return None
|
||
|
||
result = await self._db.execute(
|
||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||
)
|
||
workflow = result.scalar_one_or_none()
|
||
if not workflow:
|
||
return None
|
||
|
||
# 计算新版本号
|
||
new_version = workflow.version or 1
|
||
|
||
# 创建版本记录
|
||
version = AIWorkflowVersion(
|
||
workflow_id=workflow_id,
|
||
version=new_version,
|
||
definition=workflow.definition,
|
||
description=description,
|
||
)
|
||
self._db.add(version)
|
||
|
||
# 更新工作流状态
|
||
workflow.status = 'published'
|
||
workflow.published_version = new_version
|
||
workflow.published_at = datetime.now()
|
||
workflow.published_definition = workflow.definition
|
||
workflow.version = new_version + 1 # 草稿版本号 +1
|
||
|
||
await self._db.commit()
|
||
await self._db.refresh(workflow)
|
||
|
||
return workflow
|
||
|
||
async def rollback_workflow(self, workflow_id: str, version: int) -> Optional[AIWorkflow]:
|
||
"""
|
||
回滚到指定版本
|
||
|
||
将指定版本的定义复制到当前草稿
|
||
"""
|
||
if not self._db:
|
||
return None
|
||
|
||
result = await self._db.execute(
|
||
select(AIWorkflow).where(AIWorkflow.id == workflow_id, AIWorkflow.is_deleted == False)
|
||
)
|
||
workflow = result.scalar_one_or_none()
|
||
if not workflow:
|
||
return None
|
||
|
||
version_result = await self._db.execute(
|
||
select(AIWorkflowVersion).where(
|
||
AIWorkflowVersion.workflow_id == workflow_id,
|
||
AIWorkflowVersion.version == version,
|
||
AIWorkflowVersion.is_deleted == False
|
||
)
|
||
)
|
||
version_obj = version_result.scalar_one_or_none()
|
||
if not version_obj:
|
||
return None
|
||
|
||
# 将版本定义复制到当前草稿
|
||
workflow.definition = version_obj.definition
|
||
|
||
await self._db.commit()
|
||
await self._db.refresh(workflow)
|
||
|
||
return workflow
|
||
|
||
async def run_workflow(
|
||
self,
|
||
workflow_id: str,
|
||
inputs: Dict[str, Any],
|
||
app_id: str = None,
|
||
conversation_id: str = None,
|
||
use_draft: bool = False,
|
||
trigger_type: Optional[str] = None,
|
||
) -> AIWorkflowRun:
|
||
"""
|
||
运行工作流
|
||
|
||
Args:
|
||
workflow_id: 工作流 ID
|
||
inputs: 输入数据
|
||
app_id: 关联应用 ID
|
||
conversation_id: 关联对话 ID
|
||
use_draft: 是否使用草稿版本(编辑器调试用)
|
||
|
||
Returns:
|
||
AIWorkflowRun
|
||
"""
|
||
if not self._db:
|
||
raise ValueError("数据库会话未初始化")
|
||
|
||
workflow = await self.get_workflow(workflow_id)
|
||
if not workflow:
|
||
raise ValueError(f'工作流不存在: {workflow_id}')
|
||
|
||
if not use_draft and workflow.status != 'published':
|
||
raise ValueError('工作流未发布')
|
||
|
||
# 创建运行记录
|
||
run = self._create_workflow_run_record(
|
||
workflow,
|
||
inputs,
|
||
app_id=app_id,
|
||
conversation_id=conversation_id,
|
||
use_draft=use_draft,
|
||
trigger_type=trigger_type or 'api',
|
||
)
|
||
self._db.add(run)
|
||
await self._db.flush()
|
||
|
||
try:
|
||
# 执行工作流(异步)
|
||
result = await self._execute_workflow(
|
||
workflow, inputs, run, use_draft=use_draft
|
||
)
|
||
|
||
# 更新运行记录
|
||
run.status = 'completed'
|
||
run.outputs = dict(result.get('outputs', {}))
|
||
run.execution_log = list(result.get('logs', []))
|
||
run.total_tokens = result.get('total_tokens', 0)
|
||
run.total_steps = result.get('total_steps', 0)
|
||
run.completed_at = datetime.now()
|
||
run.elapsed_time = int((run.completed_at - run.started_at).total_seconds() * 1000)
|
||
|
||
# 更新工作流统计
|
||
workflow.run_count = (workflow.run_count or 0) + 1
|
||
workflow.success_count = (workflow.success_count or 0) + 1
|
||
|
||
await self._db.commit()
|
||
|
||
except Exception as e:
|
||
logger.exception(f'工作流执行失败: {e}')
|
||
run.status = 'failed'
|
||
run.error_message = str(e)
|
||
run.completed_at = datetime.now()
|
||
run.elapsed_time = int((run.completed_at - run.started_at).total_seconds() * 1000)
|
||
|
||
workflow.run_count = (workflow.run_count or 0) + 1
|
||
|
||
await self._db.commit()
|
||
|
||
raise
|
||
|
||
return run
|
||
|
||
async def _execute_workflow(
|
||
self,
|
||
workflow: AIWorkflow,
|
||
inputs: Dict[str, Any],
|
||
run: AIWorkflowRun,
|
||
use_draft: bool = False,
|
||
) -> Dict[str, Any]:
|
||
"""
|
||
执行工作流(异步版本)
|
||
|
||
Args:
|
||
workflow: 工作流定义
|
||
inputs: 输入数据
|
||
run: 运行记录
|
||
use_draft: 是否使用草稿版本(编辑器调试用)
|
||
|
||
Returns:
|
||
执行结果
|
||
"""
|
||
# 编辑器调试模式使用草稿版本,否则使用已发布版本
|
||
if use_draft:
|
||
definition = workflow.definition
|
||
else:
|
||
definition = workflow.published_definition or workflow.definition
|
||
nodes = definition.get('nodes', [])
|
||
edges = definition.get('edges', [])
|
||
|
||
if not nodes:
|
||
raise ValueError('工作流没有节点')
|
||
|
||
# 构建节点映射和边映射
|
||
node_map = {node['id']: node for node in nodes}
|
||
edge_map = {} # source_id -> [target_ids]
|
||
# 并行节点的分支映射:source_id -> {sourceHandle -> target_id}
|
||
parallel_edge_map = {}
|
||
for edge in edges:
|
||
source = edge.get('source', '')
|
||
target = edge.get('target', '')
|
||
source_handle = edge.get('sourceHandle', '')
|
||
if source not in edge_map:
|
||
edge_map[source] = []
|
||
edge_map[source].append(target)
|
||
# 记录 sourceHandle 映射(用于并行节点)
|
||
if source_handle:
|
||
if source not in parallel_edge_map:
|
||
parallel_edge_map[source] = {}
|
||
parallel_edge_map[source][source_handle] = target
|
||
|
||
# 找到开始节点
|
||
start_node = None
|
||
for node in nodes:
|
||
if node.get('type') == 'start':
|
||
start_node = node
|
||
break
|
||
|
||
if not start_node:
|
||
raise ValueError('找不到开始节点')
|
||
|
||
# 初始化上下文
|
||
context = NodeContext(
|
||
workflow_run_id=str(run.id),
|
||
variables=inputs.copy(),
|
||
user_input=inputs.get('user_input', ''),
|
||
user_id=str(run.user_id) if run.user_id else '',
|
||
db_session=self._db,
|
||
)
|
||
|
||
# 执行节点
|
||
logs = []
|
||
total_tokens = 0
|
||
total_steps = 0
|
||
current_node_id = start_node['id']
|
||
|
||
while current_node_id:
|
||
node_config = node_map.get(current_node_id)
|
||
if not node_config:
|
||
break
|
||
|
||
node_type = node_config.get('type', '')
|
||
node_inputs = _snapshot_node_inputs(context)
|
||
|
||
# 结束节点
|
||
if node_type == 'end':
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if node_instance:
|
||
result = node_instance.execute(context)
|
||
logs.append(_make_execution_log_entry(
|
||
node_map,
|
||
current_node_id,
|
||
node_type,
|
||
status='completed',
|
||
output=result.output,
|
||
metadata=_node_result_metadata(result),
|
||
inputs=copy.deepcopy(node_inputs),
|
||
))
|
||
break
|
||
|
||
# 创建节点实例
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if not node_instance:
|
||
logger.warning(f'未知节点类型: {node_type}')
|
||
# 跳到下一个节点
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
continue
|
||
|
||
# 执行节点
|
||
start_time = time.time()
|
||
result = await node_instance.execute_async(context)
|
||
elapsed = int((time.time() - start_time) * 1000)
|
||
target_node_id = _resolve_handoff_target(
|
||
current_node_id,
|
||
result,
|
||
edge_map,
|
||
parallel_edge_map,
|
||
node_map,
|
||
)
|
||
|
||
# 记录日志
|
||
logs.append(_make_execution_log_entry(
|
||
node_map,
|
||
current_node_id,
|
||
node_type,
|
||
status='completed' if result.success else 'failed',
|
||
output=result.output,
|
||
error=result.error,
|
||
elapsed_time=elapsed,
|
||
tokens_used=result.tokens_used,
|
||
metadata=_node_result_metadata(result),
|
||
inputs=copy.deepcopy(node_inputs),
|
||
target_node_id=target_node_id,
|
||
))
|
||
|
||
total_tokens += result.tokens_used
|
||
total_steps += 1
|
||
|
||
# 更新运行记录的当前节点并提交
|
||
run.current_node_id = current_node_id
|
||
run.execution_log = logs.copy()
|
||
if self._db:
|
||
await self._db.commit()
|
||
|
||
if not result.success:
|
||
raise ValueError(f'节点执行失败: {result.error}')
|
||
|
||
# 合并输出变量到上下文
|
||
context.variables.update(result.output_variables)
|
||
context.previous_output = result.output
|
||
|
||
# 保存节点输出到命名空间(支持 {{NodeID.key}} 格式引用)
|
||
context.variables[f'_node_{current_node_id}'] = {
|
||
'output': result.output,
|
||
**result.output_variables,
|
||
}
|
||
|
||
# 确定下一个节点
|
||
if node_type == 'parallel':
|
||
# 并行节点:执行所有分支
|
||
parallel_results = await self._execute_parallel_branches(
|
||
node_map, edge_map, parallel_edge_map, current_node_id, context, run, logs
|
||
)
|
||
# 将并行结果存入上下文
|
||
context.set_variable('_parallel_results', parallel_results['results'])
|
||
total_tokens += parallel_results['total_tokens']
|
||
total_steps += parallel_results['total_steps']
|
||
logs.extend(parallel_results['logs'])
|
||
|
||
# 找到合并节点继续执行
|
||
current_node_id = parallel_results.get('merge_node_id')
|
||
elif result.next_node_id:
|
||
# 检查是否是分支 ID(条件/意图分支节点返回的是分支 ID,需要通过边映射找到实际目标节点)
|
||
logger.info(f'节点 {current_node_id} 返回 next_node_id={result.next_node_id}')
|
||
logger.info(f'parallel_edge_map[{current_node_id}]={parallel_edge_map.get(current_node_id, {})}')
|
||
if result.next_node_id in parallel_edge_map.get(current_node_id, {}):
|
||
current_node_id = parallel_edge_map[current_node_id][result.next_node_id]
|
||
logger.info(f'通过分支映射找到目标节点: {current_node_id}')
|
||
elif result.next_node_id in node_map:
|
||
# 是实际的节点 ID
|
||
current_node_id = result.next_node_id
|
||
logger.info(f'直接使用节点 ID: {current_node_id}')
|
||
else:
|
||
# 分支 ID 未找到对应边,尝试按普通边查找
|
||
logger.warning(f'分支 ID {result.next_node_id} 未找到对应边,回退到普通边查找')
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
else:
|
||
# 按边查找下一个节点
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
|
||
return {
|
||
'outputs': context.variables,
|
||
'logs': logs,
|
||
'total_tokens': total_tokens,
|
||
'total_steps': total_steps,
|
||
}
|
||
|
||
async def _execute_parallel_branches(
|
||
self,
|
||
node_map: Dict[str, Any],
|
||
edge_map: Dict[str, List[str]],
|
||
parallel_edge_map: Dict[str, Dict[str, str]],
|
||
parallel_node_id: str,
|
||
context: NodeContext,
|
||
run: AIWorkflowRun,
|
||
parent_logs: List[Dict],
|
||
) -> Dict[str, Any]:
|
||
"""
|
||
执行并行分支(异步版本)
|
||
|
||
Args:
|
||
node_map: 节点映射
|
||
edge_map: 边映射
|
||
parallel_edge_map: 并行节点的分支映射 {source_id: {sourceHandle: target_id}}
|
||
parallel_node_id: 并行节点 ID
|
||
context: 执行上下文
|
||
run: 运行记录
|
||
parent_logs: 父级日志列表
|
||
|
||
Returns:
|
||
并行执行结果
|
||
"""
|
||
import asyncio
|
||
|
||
# 获取并行节点的分支映射(sourceHandle -> target_id)
|
||
branch_mapping = parallel_edge_map.get(parallel_node_id, {})
|
||
|
||
# 如果没有分支映射,回退到普通边映射
|
||
if not branch_mapping:
|
||
branch_start_nodes = edge_map.get(parallel_node_id, [])
|
||
# 创建默认的分支映射
|
||
branch_mapping = {f'branch_{i}': node_id for i, node_id in enumerate(branch_start_nodes)}
|
||
|
||
branch_start_nodes = list(branch_mapping.values())
|
||
|
||
if not branch_start_nodes:
|
||
return {
|
||
'results': {},
|
||
'logs': [],
|
||
'total_tokens': 0,
|
||
'total_steps': 0,
|
||
'merge_node_id': None,
|
||
}
|
||
|
||
# 找到合并节点(所有分支最终汇聚的节点)
|
||
merge_node_id = self._find_merge_node(node_map, edge_map, branch_start_nodes)
|
||
|
||
# 并行执行各分支
|
||
branch_results = {}
|
||
all_logs = []
|
||
total_tokens = 0
|
||
total_steps = 0
|
||
branch_labels = _parallel_branch_labels(node_map, parallel_node_id)
|
||
|
||
async def execute_branch(branch_id: str, branch_start_id: str, branch_context: NodeContext) -> Dict:
|
||
"""执行单个分支(异步)"""
|
||
branch_logs = []
|
||
branch_tokens = 0
|
||
branch_steps = 0
|
||
current_id = branch_start_id
|
||
branch_output = None
|
||
|
||
while current_id and current_id != merge_node_id:
|
||
node_config = node_map.get(current_id)
|
||
if not node_config:
|
||
break
|
||
|
||
node_type = node_config.get('type', '')
|
||
|
||
# 跳过结束节点和合并节点
|
||
if node_type in ('end', 'merge'):
|
||
break
|
||
|
||
# 创建节点实例
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if not node_instance:
|
||
next_nodes = edge_map.get(current_id, [])
|
||
current_id = next_nodes[0] if next_nodes else None
|
||
continue
|
||
|
||
node_inputs = _snapshot_node_inputs(branch_context)
|
||
|
||
# 执行节点(异步)。并行分支不能共享同一个 AsyncSession,否则模型/Provider 查询会并发抢占连接。
|
||
start_time = time.time()
|
||
if context.db_session is not None:
|
||
from app.database import AsyncSessionLocal
|
||
|
||
async with AsyncSessionLocal() as branch_db:
|
||
branch_context.db_session = branch_db
|
||
result = await node_instance.execute_async(branch_context)
|
||
else:
|
||
result = await node_instance.execute_async(branch_context)
|
||
elapsed = int((time.time() - start_time) * 1000)
|
||
target_node_id = _resolve_handoff_target(
|
||
current_id,
|
||
result,
|
||
edge_map,
|
||
parallel_edge_map,
|
||
node_map,
|
||
)
|
||
|
||
branch_logs.append(_make_execution_log_entry(
|
||
node_map,
|
||
current_id,
|
||
node_type,
|
||
status='completed' if result.success else 'failed',
|
||
output=result.output,
|
||
error=result.error,
|
||
elapsed_time=elapsed,
|
||
tokens_used=result.tokens_used,
|
||
metadata=_node_result_metadata(result),
|
||
inputs=copy.deepcopy(node_inputs),
|
||
branch=branch_id,
|
||
branch_label=branch_labels.get(branch_id, branch_id),
|
||
target_node_id=target_node_id,
|
||
))
|
||
|
||
branch_tokens += result.tokens_used
|
||
branch_steps += 1
|
||
|
||
if not result.success:
|
||
return {
|
||
'branch_id': branch_id,
|
||
'success': False,
|
||
'output': None,
|
||
'error': result.error,
|
||
'logs': branch_logs,
|
||
'tokens': branch_tokens,
|
||
'steps': branch_steps,
|
||
}
|
||
|
||
# 更新分支上下文
|
||
branch_context.variables.update(result.output_variables)
|
||
branch_context.previous_output = result.output
|
||
branch_output = result.output
|
||
|
||
# 支持分支内后续节点使用 {{node_id.output_variable}} 引用当前节点输出。
|
||
branch_context.variables[f'_node_{current_id}'] = {
|
||
'output': result.output,
|
||
**result.output_variables,
|
||
}
|
||
|
||
# 下一个节点
|
||
if result.next_node_id:
|
||
current_id = result.next_node_id
|
||
else:
|
||
next_nodes = edge_map.get(current_id, [])
|
||
# 过滤掉合并节点
|
||
next_nodes = [n for n in next_nodes if n != merge_node_id]
|
||
current_id = next_nodes[0] if next_nodes else None
|
||
|
||
return {
|
||
'branch_id': branch_id,
|
||
'success': True,
|
||
'output': branch_output,
|
||
'logs': branch_logs,
|
||
'tokens': branch_tokens,
|
||
'steps': branch_steps,
|
||
'variables': branch_context.variables,
|
||
}
|
||
|
||
# 创建反向映射:target_node_id -> branch_id (sourceHandle)
|
||
target_to_branch = {v: k for k, v in branch_mapping.items()}
|
||
|
||
# 创建所有分支的任务
|
||
tasks = []
|
||
for target_node_id in branch_start_nodes:
|
||
# 为每个分支创建独立的上下文副本
|
||
branch_context = NodeContext(
|
||
workflow_run_id=context.workflow_run_id,
|
||
variables=copy.deepcopy(context.variables),
|
||
user_input=context.user_input,
|
||
user_id=context.user_id,
|
||
conversation_history=context.conversation_history.copy(),
|
||
previous_output=context.previous_output,
|
||
db_session=context.db_session,
|
||
)
|
||
# 获取分支 ID(sourceHandle)
|
||
branch_id = target_to_branch.get(target_node_id, target_node_id)
|
||
tasks.append(execute_branch(branch_id, target_node_id, branch_context))
|
||
|
||
# 并行执行所有分支
|
||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||
|
||
for result in results:
|
||
if isinstance(result, Exception):
|
||
logger.exception(f'分支执行失败: {result}')
|
||
all_logs.append({
|
||
'node_id': 'unknown',
|
||
'node_type': 'branch',
|
||
'status': 'failed',
|
||
'error': str(result),
|
||
})
|
||
else:
|
||
branch_id = result.get('branch_id')
|
||
branch_results[branch_id] = result.get('output')
|
||
all_logs.extend(result.get('logs', []))
|
||
total_tokens += result.get('tokens', 0)
|
||
total_steps += result.get('steps', 0)
|
||
|
||
# 合并分支变量到主上下文
|
||
if result.get('variables'):
|
||
for key, value in result['variables'].items():
|
||
if not key.startswith('_'): # 不合并内部变量
|
||
context.variables[f'{branch_id}_{key}'] = value
|
||
|
||
return {
|
||
'results': branch_results,
|
||
'logs': all_logs,
|
||
'total_tokens': total_tokens,
|
||
'total_steps': total_steps,
|
||
'merge_node_id': merge_node_id,
|
||
}
|
||
|
||
async def _execute_parallel_branches_stream(
|
||
self,
|
||
node_map: Dict[str, Any],
|
||
edge_map: Dict[str, List[str]],
|
||
parallel_edge_map: Dict[str, Dict[str, str]],
|
||
parallel_node_id: str,
|
||
context: NodeContext,
|
||
run: AIWorkflowRun,
|
||
parent_logs: List[Dict],
|
||
):
|
||
"""
|
||
流式执行并行分支,yield 每个节点的执行事件(异步版本)
|
||
|
||
Args:
|
||
node_map: 节点映射
|
||
edge_map: 边映射
|
||
parallel_edge_map: 并行节点的分支映射
|
||
parallel_node_id: 并行节点 ID
|
||
context: 执行上下文
|
||
run: 运行记录
|
||
parent_logs: 父级日志列表
|
||
|
||
Yields:
|
||
dict: 节点执行事件
|
||
"""
|
||
import asyncio
|
||
|
||
# 获取并行节点的分支映射
|
||
branch_mapping = parallel_edge_map.get(parallel_node_id, {})
|
||
|
||
if not branch_mapping:
|
||
branch_start_nodes = edge_map.get(parallel_node_id, [])
|
||
branch_mapping = {f'branch_{i}': node_id for i, node_id in enumerate(branch_start_nodes)}
|
||
|
||
branch_start_nodes = list(branch_mapping.values())
|
||
|
||
if not branch_start_nodes:
|
||
yield {
|
||
'type': '_parallel_result',
|
||
'data': {
|
||
'results': {},
|
||
'logs': [],
|
||
'total_tokens': 0,
|
||
'total_steps': 0,
|
||
'merge_node_id': None,
|
||
}
|
||
}
|
||
return
|
||
|
||
merge_node_id = self._find_merge_node(node_map, edge_map, branch_start_nodes)
|
||
|
||
branch_results = {}
|
||
all_logs = []
|
||
total_tokens = 0
|
||
total_steps = 0
|
||
branch_labels = _parallel_branch_labels(node_map, parallel_node_id)
|
||
|
||
target_to_branch = {v: k for k, v in branch_mapping.items()}
|
||
|
||
# 顺序执行各分支(为了能够yield事件)
|
||
for target_node_id in branch_start_nodes:
|
||
branch_id = target_to_branch.get(target_node_id, target_node_id)
|
||
branch_context = NodeContext(
|
||
workflow_run_id=context.workflow_run_id,
|
||
variables=copy.deepcopy(context.variables),
|
||
user_input=context.user_input,
|
||
user_id=context.user_id,
|
||
conversation_history=context.conversation_history.copy(),
|
||
previous_output=context.previous_output,
|
||
db_session=context.db_session,
|
||
)
|
||
|
||
branch_logs = []
|
||
branch_tokens = 0
|
||
branch_steps = 0
|
||
current_id = target_node_id
|
||
branch_output = None
|
||
|
||
while current_id and current_id != merge_node_id:
|
||
node_config = node_map.get(current_id)
|
||
if not node_config:
|
||
break
|
||
|
||
node_type = node_config.get('type', '')
|
||
node_label = node_config.get('data', {}).get('label', current_id)
|
||
|
||
if node_type in ('end', 'merge'):
|
||
break
|
||
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if not node_instance:
|
||
next_nodes = edge_map.get(current_id, [])
|
||
current_id = next_nodes[0] if next_nodes else None
|
||
continue
|
||
|
||
# 发送节点开始事件
|
||
yield _merge_event_collaboration({
|
||
'type': 'node_start',
|
||
'node_id': current_id,
|
||
'node_type': node_type,
|
||
'node_label': node_label,
|
||
'branch_id': branch_id,
|
||
'branch_label': branch_labels.get(branch_id, branch_id),
|
||
'timestamp': datetime.now().isoformat(),
|
||
'inputs': {
|
||
'variables': dict(branch_context.variables),
|
||
'previous_output': branch_context.previous_output,
|
||
},
|
||
}, node_map, current_id, node_type)
|
||
|
||
# 执行节点(异步)
|
||
start_time = time.time()
|
||
result = await node_instance.execute_async(branch_context)
|
||
elapsed = int((time.time() - start_time) * 1000)
|
||
|
||
branch_logs.append(_make_execution_log_entry(
|
||
node_map,
|
||
current_id,
|
||
node_type,
|
||
status='completed' if result.success else 'failed',
|
||
output=result.output,
|
||
error=result.error,
|
||
elapsed_time=elapsed,
|
||
tokens_used=result.tokens_used,
|
||
metadata=_node_result_metadata(result),
|
||
branch=branch_id,
|
||
branch_label=branch_labels.get(branch_id, branch_id),
|
||
))
|
||
|
||
branch_tokens += result.tokens_used
|
||
branch_steps += 1
|
||
|
||
# 发送节点完成事件
|
||
yield _merge_event_collaboration(_apply_node_result_warnings({
|
||
'type': 'node_complete',
|
||
'node_id': current_id,
|
||
'node_type': node_type,
|
||
'branch_id': branch_id,
|
||
'branch_label': branch_labels.get(branch_id, branch_id),
|
||
'status': 'success' if result.success else 'failed',
|
||
'elapsed_time': elapsed,
|
||
'tokens_used': result.tokens_used,
|
||
'error_message': result.error if not result.success else None,
|
||
'outputs': {
|
||
'output': result.output,
|
||
'output_variables': result.metadata.get('frontend_output_variables', result.output_variables),
|
||
},
|
||
'target_node_id': target_node_id,
|
||
}, result), node_map, current_id, node_type, result)
|
||
|
||
if not result.success:
|
||
break
|
||
|
||
branch_context.variables.update(result.output_variables)
|
||
branch_context.previous_output = result.output
|
||
branch_output = result.output
|
||
|
||
# 保存节点输出到命名空间
|
||
branch_context.variables[f'_node_{current_id}'] = {
|
||
'output': result.output,
|
||
**result.output_variables,
|
||
}
|
||
|
||
if result.next_node_id:
|
||
current_id = result.next_node_id
|
||
else:
|
||
next_nodes = edge_map.get(current_id, [])
|
||
next_nodes = [n for n in next_nodes if n != merge_node_id]
|
||
current_id = next_nodes[0] if next_nodes else None
|
||
|
||
# 收集分支结果
|
||
branch_results[branch_id] = branch_output
|
||
all_logs.extend(branch_logs)
|
||
total_tokens += branch_tokens
|
||
total_steps += branch_steps
|
||
|
||
# 合并分支变量到主上下文
|
||
for key, value in branch_context.variables.items():
|
||
if not key.startswith('_'):
|
||
context.variables[f'{branch_id}_{key}'] = value
|
||
|
||
# 返回最终结果
|
||
yield {
|
||
'type': '_parallel_result',
|
||
'data': {
|
||
'results': branch_results,
|
||
'logs': all_logs,
|
||
'total_tokens': total_tokens,
|
||
'total_steps': total_steps,
|
||
'merge_node_id': merge_node_id,
|
||
}
|
||
}
|
||
|
||
def _find_merge_node(
|
||
self,
|
||
node_map: Dict[str, Any],
|
||
edge_map: Dict[str, List[str]],
|
||
branch_start_nodes: List[str],
|
||
) -> Optional[str]:
|
||
"""
|
||
查找合并节点
|
||
|
||
遍历所有分支,找到它们共同汇聚的合并节点
|
||
"""
|
||
# 简单实现:查找类型为 merge 的节点
|
||
for node_id, node_config in node_map.items():
|
||
if node_config.get('type') == 'merge':
|
||
return node_id
|
||
|
||
# 如果没有显式的合并节点,尝试找到所有分支汇聚的节点
|
||
# 这里使用简单的方法:找到第一个被多个节点指向的节点
|
||
incoming_count = {}
|
||
for source, targets in edge_map.items():
|
||
for target in targets:
|
||
incoming_count[target] = incoming_count.get(target, 0) + 1
|
||
|
||
for node_id, count in incoming_count.items():
|
||
if count >= len(branch_start_nodes):
|
||
return node_id
|
||
|
||
return None
|
||
|
||
def _get_child_nodes(
|
||
self,
|
||
node_map: Dict[str, Any],
|
||
parent_node_id: str,
|
||
) -> List[str]:
|
||
"""
|
||
获取循环节点的子节点列表
|
||
|
||
子节点通过 parentNode 属性关联到循环节点
|
||
"""
|
||
child_nodes = []
|
||
for node_id, node_config in node_map.items():
|
||
# 检查节点配置中的 parentNode 属性
|
||
parent_node = node_config.get('parentNode')
|
||
if parent_node == parent_node_id:
|
||
child_nodes.append(node_id)
|
||
logger.info(f'[Loop] Found child node: {node_id} with parentNode={parent_node}')
|
||
|
||
logger.info(f'[Loop] Parent node {parent_node_id} has {len(child_nodes)} child nodes: {child_nodes}')
|
||
return child_nodes
|
||
|
||
async def _execute_loop(
|
||
self,
|
||
node_map: Dict[str, Any],
|
||
edge_map: Dict[str, List[str]],
|
||
parallel_edge_map: Dict[str, Dict[str, str]],
|
||
loop_node_id: str,
|
||
loop_config: Dict[str, Any],
|
||
context: NodeContext,
|
||
run: AIWorkflowRun,
|
||
parent_logs: List[Dict],
|
||
):
|
||
"""
|
||
执行循环体(单节点设计,异步版本)
|
||
|
||
循环体的子节点通过 parentNode 属性关联到循环节点,
|
||
或者通过 childNodes 配置指定
|
||
|
||
Args:
|
||
node_map: 节点映射
|
||
edge_map: 边映射
|
||
parallel_edge_map: 并行边映射
|
||
loop_node_id: 循环节点 ID
|
||
loop_config: 循环配置
|
||
context: 执行上下文
|
||
run: 运行记录
|
||
parent_logs: 父级日志列表
|
||
|
||
Yields:
|
||
循环执行事件
|
||
"""
|
||
import asyncio
|
||
loop_mode = context.get_variable('_loop_mode', 'for_each')
|
||
max_iterations = loop_config.get('max_iterations', 100)
|
||
item_var_name = loop_config.get('item_variable_name', 'item')
|
||
index_var_name = loop_config.get('index_variable_name', 'index')
|
||
|
||
# 获取循环体子节点
|
||
# 方式1: 通过 parentNode 属性关联
|
||
child_node_ids = self._get_child_nodes(node_map, loop_node_id)
|
||
|
||
# 方式2: 通过 childNodes 配置指定
|
||
if not child_node_ids:
|
||
child_node_ids = loop_config.get('childNodes', [])
|
||
|
||
if not child_node_ids:
|
||
logger.warning(f'循环节点 {loop_node_id} 没有子节点')
|
||
return
|
||
|
||
# 构建子节点的边映射(只包含子节点之间的边)
|
||
child_edge_map = {}
|
||
for source_id in child_node_ids:
|
||
targets = edge_map.get(source_id, [])
|
||
# 只保留指向其他子节点的边
|
||
child_targets = [t for t in targets if t in child_node_ids]
|
||
if child_targets:
|
||
child_edge_map[source_id] = child_targets
|
||
|
||
# 找到循环体的起始节点(没有入边的子节点)
|
||
nodes_with_incoming = set()
|
||
for targets in child_edge_map.values():
|
||
nodes_with_incoming.update(targets)
|
||
|
||
start_nodes = [n for n in child_node_ids if n not in nodes_with_incoming]
|
||
if not start_nodes:
|
||
# 如果所有节点都有入边,取第一个子节点
|
||
start_nodes = [child_node_ids[0]]
|
||
|
||
loop_body_start_id = start_nodes[0]
|
||
|
||
# 从当前循环索引开始(支持恢复执行)
|
||
iteration = context.get_variable('_loop_index', 0)
|
||
logger.info(f'[Loop] Starting loop execution, mode={loop_mode}, max_iterations={max_iterations}, start_iteration={iteration}')
|
||
logger.info(f'[Loop] _loop_items={context.get_variable("_loop_items", [])}')
|
||
logger.info(f'[Loop] _loop_index={context.get_variable("_loop_index", 0)}')
|
||
|
||
while iteration < max_iterations:
|
||
# 检查是否继续循环
|
||
if loop_mode == 'for_each':
|
||
items = context.get_variable('_loop_items', [])
|
||
current_index = context.get_variable('_loop_index', 0)
|
||
logger.info(f'[Loop] Iteration {iteration}: items={items}, current_index={current_index}')
|
||
if current_index >= len(items):
|
||
logger.info(f'[Loop] Breaking loop: current_index({current_index}) >= len(items)({len(items)})')
|
||
break
|
||
# 设置当前循环项和索引
|
||
current_item = items[current_index]
|
||
context.set_variable(item_var_name, current_item)
|
||
context.set_variable(index_var_name, current_index)
|
||
context.set_variable('_current_item', current_item)
|
||
context.set_variable('_current_index', current_index)
|
||
# 同时存储到循环节点的命名空间,支持 {{loop_node_id.item.xxx}} 格式
|
||
context.variables[f'_node_{loop_node_id}'] = {
|
||
item_var_name: current_item,
|
||
index_var_name: current_index,
|
||
'total': len(items),
|
||
}
|
||
elif loop_mode == 'while':
|
||
# while 模式:需要重新评估条件
|
||
from ai_platform.nodes.builtin.loop_node import LoopNode
|
||
loop_node_instance = LoopNode(loop_config)
|
||
if not loop_node_instance._evaluate_condition(context):
|
||
break
|
||
context.set_variable(index_var_name, iteration)
|
||
|
||
# 发送迭代开始事件
|
||
total_iterations = len(items) if loop_mode == 'for_each' else None
|
||
yield {
|
||
'type': 'loop_iteration_start',
|
||
'node_id': loop_node_id,
|
||
'iteration': iteration,
|
||
'total': total_iterations,
|
||
'item': context.get_variable(item_var_name) if loop_mode == 'for_each' else None,
|
||
}
|
||
|
||
# 执行循环体
|
||
iteration_logs = []
|
||
iteration_tokens = 0
|
||
iteration_steps = 0
|
||
current_id = loop_body_start_id
|
||
iteration_output = None
|
||
executed_nodes = set()
|
||
|
||
while current_id and current_id in child_node_ids:
|
||
# 防止无限循环
|
||
if current_id in executed_nodes:
|
||
break
|
||
executed_nodes.add(current_id)
|
||
|
||
node_config = node_map.get(current_id)
|
||
if not node_config:
|
||
break
|
||
|
||
node_type = node_config.get('type', '')
|
||
|
||
# 创建节点实例
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if not node_instance:
|
||
next_nodes = child_edge_map.get(current_id, [])
|
||
current_id = next_nodes[0] if next_nodes else None
|
||
continue
|
||
|
||
# 发送节点开始事件
|
||
yield _merge_event_collaboration({
|
||
'type': 'node_start',
|
||
'node_id': current_id,
|
||
'node_type': node_type,
|
||
'node_label': node_config.get('data', {}).get('label', ''),
|
||
'loop_iteration': iteration,
|
||
'inputs': {
|
||
'previous_output': context.previous_output,
|
||
'variables': {
|
||
'_current_item': context.get_variable('_current_item'),
|
||
'_current_index': context.get_variable('_current_index'),
|
||
'_loop_total': context.get_variable('_loop_total'),
|
||
}
|
||
},
|
||
}, node_map, current_id, node_type)
|
||
|
||
# 执行节点(异步)
|
||
start_time = time.time()
|
||
logger.info(f'[Loop] Executing child node: {current_id}, type={node_type}')
|
||
result = await node_instance.execute_async(context)
|
||
elapsed = int((time.time() - start_time) * 1000)
|
||
logger.info(f'[Loop] Child node result: success={result.success}, output={result.output}, error={result.error}')
|
||
|
||
iteration_logs.append(_make_execution_log_entry(
|
||
node_map,
|
||
current_id,
|
||
node_type,
|
||
status='completed' if result.success else 'failed',
|
||
output=result.output,
|
||
error=result.error,
|
||
elapsed_time=elapsed,
|
||
tokens_used=result.tokens_used,
|
||
metadata=_node_result_metadata(result),
|
||
loop_iteration=iteration,
|
||
))
|
||
|
||
iteration_tokens += result.tokens_used
|
||
iteration_steps += 1
|
||
|
||
# 发送节点执行完成事件
|
||
yield _merge_event_collaboration(_apply_node_result_warnings({
|
||
'type': 'node_complete',
|
||
'node_id': current_id,
|
||
'node_type': node_type,
|
||
'status': 'success' if result.success else 'failed',
|
||
'elapsed_time': elapsed,
|
||
'tokens_used': result.tokens_used,
|
||
'error_message': result.error,
|
||
'loop_iteration': iteration,
|
||
'outputs': {
|
||
'output': result.output,
|
||
'output_variables': result.metadata.get('frontend_output_variables', result.output_variables),
|
||
},
|
||
'inputs': {
|
||
'previous_output': context.previous_output,
|
||
'variables': {
|
||
'_current_item': context.get_variable('_current_item'),
|
||
'_current_index': context.get_variable('_current_index'),
|
||
'_loop_total': context.get_variable('_loop_total'),
|
||
}
|
||
},
|
||
}, result), node_map, current_id, node_type, result)
|
||
|
||
# 发送节点产生的事件(如消息事件)
|
||
if result.events:
|
||
for event in result.events:
|
||
yield {
|
||
'type': 'node_event',
|
||
'node_id': current_id,
|
||
'event': event,
|
||
'loop_iteration': iteration,
|
||
}
|
||
|
||
# 检查是否需要等待用户输入(设计预览节点等)
|
||
if result.waiting_for_input:
|
||
# 记录日志
|
||
iteration_logs.append(_make_execution_log_entry(
|
||
node_map,
|
||
current_id,
|
||
node_type,
|
||
status='waiting',
|
||
elapsed_time=elapsed,
|
||
loop_iteration=iteration,
|
||
))
|
||
|
||
# 更新运行记录为等待状态
|
||
run.status = 'waiting'
|
||
run.current_node_id = current_id
|
||
run.execution_log = parent_logs + iteration_logs
|
||
run.outputs = dict(context.variables)
|
||
run.waiting_config = dict(result.waiting_config) if result.waiting_config else {}
|
||
# 保存循环状态以便恢复
|
||
run.waiting_config['_loop_state'] = {
|
||
'loop_node_id': loop_node_id,
|
||
'iteration': iteration,
|
||
'current_index': context.get_variable('_loop_index', 0),
|
||
}
|
||
if self._db:
|
||
await self._db.commit()
|
||
|
||
# 发送等待事件
|
||
yield _merge_event_collaboration({
|
||
'type': 'waiting_input',
|
||
'node_id': current_id,
|
||
'node_type': node_type,
|
||
'config': result.waiting_config,
|
||
'loop_iteration': iteration,
|
||
}, node_map, current_id, node_type, result)
|
||
# 暂停循环执行,等待用户输入
|
||
return
|
||
|
||
if not result.success:
|
||
logger.error(f'[Loop] 子节点执行失败,发送 loop_iteration_error 事件: {result.error}')
|
||
yield {
|
||
'type': 'loop_iteration_error',
|
||
'node_id': loop_node_id,
|
||
'iteration': iteration,
|
||
'error': result.error,
|
||
}
|
||
# 循环中断
|
||
return
|
||
|
||
# 更新上下文
|
||
context.variables.update(result.output_variables)
|
||
context.previous_output = result.output
|
||
iteration_output = result.output
|
||
|
||
# 保存节点输出到命名空间
|
||
context.variables[f'_node_{current_id}'] = {
|
||
'output': result.output,
|
||
**result.output_variables,
|
||
}
|
||
|
||
# 下一个节点(只在子节点范围内)
|
||
if result.next_node_id and result.next_node_id in child_node_ids:
|
||
current_id = result.next_node_id
|
||
else:
|
||
next_nodes = child_edge_map.get(current_id, [])
|
||
current_id = next_nodes[0] if next_nodes else None
|
||
|
||
# 收集本次迭代结果
|
||
loop_results = context.get_variable('_loop_results', [])
|
||
loop_results.append(iteration_output)
|
||
context.set_variable('_loop_results', loop_results)
|
||
|
||
# 更新循环索引
|
||
context.set_variable('_loop_index', context.get_variable('_loop_index', 0) + 1)
|
||
|
||
# 添加日志
|
||
parent_logs.extend(iteration_logs)
|
||
|
||
# 发送迭代完成事件
|
||
yield {
|
||
'type': 'loop_iteration_complete',
|
||
'node_id': loop_node_id,
|
||
'iteration': iteration,
|
||
'total': total_iterations,
|
||
'output': iteration_output,
|
||
'tokens_used': iteration_tokens,
|
||
'steps': iteration_steps,
|
||
}
|
||
|
||
iteration += 1
|
||
|
||
# 清理循环变量
|
||
context.variables.pop('_loop_items', None)
|
||
context.variables.pop('_loop_index', None)
|
||
context.variables.pop('_loop_total', None)
|
||
context.variables.pop('_loop_mode', None)
|
||
context.variables.pop('_current_item', None)
|
||
context.variables.pop('_current_index', None)
|
||
|
||
async def run_workflow_stream_async(
|
||
self,
|
||
workflow_id: str,
|
||
inputs: Dict[str, Any],
|
||
app_id: str = None,
|
||
conversation_id: str = None,
|
||
conversation_history: List[Dict[str, str]] = None,
|
||
use_draft: bool = False,
|
||
trigger_type: Optional[str] = None,
|
||
):
|
||
"""
|
||
异步流式运行工作流
|
||
|
||
Args:
|
||
workflow_id: 工作流 ID
|
||
inputs: 输入变量
|
||
app_id: 应用 ID
|
||
conversation_id: 对话 ID
|
||
conversation_history: 对话历史
|
||
use_draft: 是否使用草稿版本
|
||
|
||
Yields:
|
||
dict: 节点执行状态事件
|
||
"""
|
||
if not self._db:
|
||
yield {'type': 'error', 'message': '数据库会话未初始化'}
|
||
return
|
||
|
||
workflow = await self.get_workflow(workflow_id)
|
||
if not workflow:
|
||
yield {'type': 'error', 'message': f'工作流不存在: {workflow_id}'}
|
||
return
|
||
|
||
# 编辑器调试模式不检查发布状态
|
||
if not use_draft and workflow.status != 'published':
|
||
yield {'type': 'error', 'message': '工作流未发布'}
|
||
return
|
||
|
||
# 创建运行记录
|
||
run = self._create_workflow_run_record(
|
||
workflow,
|
||
inputs,
|
||
app_id=app_id,
|
||
conversation_id=conversation_id,
|
||
use_draft=use_draft,
|
||
trigger_type=trigger_type,
|
||
)
|
||
self._db.add(run)
|
||
await self._db.flush()
|
||
|
||
try:
|
||
# 使用异步生成器执行工作流
|
||
async for event in self._execute_workflow_stream(workflow, inputs, run, conversation_history, use_draft):
|
||
yield event
|
||
|
||
# 在关键事件后提交数据库状态
|
||
event_type = event.get('type')
|
||
if event_type in ('node_complete', 'waiting_input', 'complete', 'error'):
|
||
await self._db.commit()
|
||
|
||
# 更新运行记录
|
||
if run.status == 'waiting':
|
||
# 等待状态,已在循环中提交
|
||
pass
|
||
elif run.status != 'completed':
|
||
run.status = 'completed'
|
||
run.completed_at = datetime.now()
|
||
run.elapsed_time = int((run.completed_at - run.started_at).total_seconds() * 1000)
|
||
|
||
# 更新工作流统计
|
||
workflow.run_count = (workflow.run_count or 0) + 1
|
||
workflow.success_count = (workflow.success_count or 0) + 1
|
||
|
||
await self._db.commit()
|
||
|
||
except Exception as e:
|
||
logger.exception(f'工作流执行失败: {e}')
|
||
run.status = 'failed'
|
||
run.error_message = str(e)
|
||
run.completed_at = datetime.now()
|
||
run.elapsed_time = int((run.completed_at - run.started_at).total_seconds() * 1000)
|
||
|
||
workflow.run_count = (workflow.run_count or 0) + 1
|
||
|
||
await self._db.commit()
|
||
|
||
# 发送错误事件给前端
|
||
yield {
|
||
'type': 'error',
|
||
'message': str(e),
|
||
'error_message': str(e),
|
||
'run_id': run.id,
|
||
'workflow_id': workflow.id,
|
||
}
|
||
|
||
async def _execute_workflow_stream(
|
||
self,
|
||
workflow: AIWorkflow,
|
||
inputs: Dict[str, Any],
|
||
run: AIWorkflowRun,
|
||
conversation_history: List[Dict[str, str]] = None,
|
||
use_draft: bool = False,
|
||
):
|
||
"""
|
||
流式执行工作流,逐节点 yield 状态(异步版本)
|
||
|
||
Args:
|
||
workflow: 工作流对象
|
||
inputs: 输入变量
|
||
run: 工作流运行记录
|
||
conversation_history: 对话历史(用于 LLM 节点的上下文记忆)
|
||
use_draft: 是否使用草稿版本(编辑器调试用)
|
||
"""
|
||
import asyncio
|
||
# 编辑器调试模式使用草稿版本,否则使用已发布版本
|
||
if use_draft:
|
||
definition = workflow.definition
|
||
else:
|
||
definition = workflow.published_definition or workflow.definition
|
||
nodes = definition.get('nodes', [])
|
||
edges = definition.get('edges', [])
|
||
|
||
if not nodes:
|
||
raise ValueError('工作流没有节点')
|
||
|
||
# 构建节点映射和边映射
|
||
node_map = {node['id']: node for node in nodes}
|
||
edge_map = {}
|
||
# 并行节点的分支映射:source_id -> {sourceHandle -> target_id}
|
||
parallel_edge_map = {}
|
||
|
||
# 调试:打印所有节点的 parentNode 属性
|
||
for node in nodes:
|
||
if node.get('parentNode'):
|
||
logger.info(f'[流式] 节点 {node["id"]} 的 parentNode={node.get("parentNode")}')
|
||
|
||
logger.info(f'[流式] 边列表: {edges}')
|
||
for edge in edges:
|
||
source = edge.get('source', '')
|
||
target = edge.get('target', '')
|
||
source_handle = edge.get('sourceHandle', '')
|
||
logger.info(f'[流式] 边: source={source}, target={target}, sourceHandle={source_handle}')
|
||
if source not in edge_map:
|
||
edge_map[source] = []
|
||
edge_map[source].append(target)
|
||
# 记录 sourceHandle 映射(用于并行节点)
|
||
if source_handle:
|
||
if source not in parallel_edge_map:
|
||
parallel_edge_map[source] = {}
|
||
parallel_edge_map[source][source_handle] = target
|
||
|
||
# 找到开始节点
|
||
start_node = None
|
||
for node in nodes:
|
||
if node.get('type') == 'start':
|
||
start_node = node
|
||
break
|
||
|
||
if not start_node:
|
||
raise ValueError('找不到开始节点')
|
||
|
||
# 初始化上下文
|
||
context = NodeContext(
|
||
workflow_run_id=str(run.id),
|
||
variables=inputs.copy(),
|
||
user_input=inputs.get('user_input', ''),
|
||
user_id=str(run.user_id) if run.user_id else '',
|
||
conversation_history=conversation_history or [],
|
||
db_session=self._db,
|
||
)
|
||
|
||
logs = []
|
||
total_tokens = 0
|
||
total_steps = 0
|
||
current_node_id = start_node['id']
|
||
|
||
# 发送工作流开始事件(包含 run_id,用于对话流恢复)
|
||
yield {
|
||
'type': 'start',
|
||
'run_id': str(run.id),
|
||
'workflow_id': str(workflow.id),
|
||
'workflow_name': workflow.name,
|
||
}
|
||
|
||
while current_node_id:
|
||
node_config = node_map.get(current_node_id)
|
||
if not node_config:
|
||
break
|
||
|
||
node_type = node_config.get('type', '')
|
||
node_label = node_config.get('data', {}).get('label', current_node_id)
|
||
|
||
# 获取节点输入(当前上下文变量的快照)
|
||
node_inputs = {
|
||
'variables': dict(context.variables),
|
||
'previous_output': context.previous_output,
|
||
}
|
||
|
||
# 发送节点开始事件(包含时间戳和输入数据)
|
||
yield _merge_event_collaboration({
|
||
'type': 'node_start',
|
||
'node_id': current_node_id,
|
||
'node_type': node_type,
|
||
'node_label': node_label,
|
||
'timestamp': datetime.now().isoformat(),
|
||
'inputs': node_inputs,
|
||
}, node_map, current_node_id, node_type)
|
||
|
||
# 结束节点
|
||
if node_type == 'end':
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if node_instance:
|
||
result = node_instance.execute(context)
|
||
logs.append(_make_execution_log_entry(
|
||
node_map,
|
||
current_node_id,
|
||
node_type,
|
||
status='completed',
|
||
output=result.output,
|
||
metadata=_node_result_metadata(result),
|
||
inputs=copy.deepcopy(node_inputs),
|
||
))
|
||
yield _merge_event_collaboration({
|
||
'type': 'node_complete',
|
||
'node_id': current_node_id,
|
||
'node_type': node_type,
|
||
'status': 'success',
|
||
'elapsed_time': 0,
|
||
'tokens_used': 0,
|
||
}, node_map, current_node_id, node_type, result)
|
||
|
||
# 保存结束节点输出和上下文变量到运行记录
|
||
run.execution_log = logs.copy()
|
||
# 只有当结束节点有输出时才包含 output 字段
|
||
if result.output is not None:
|
||
run.outputs = {
|
||
**context.variables,
|
||
'output': result.output, # 结束节点的输出内容
|
||
}
|
||
else:
|
||
run.outputs = dict(context.variables)
|
||
|
||
# 更新运行状态为完成
|
||
run.status = 'completed'
|
||
run.completed_at = datetime.now()
|
||
run.elapsed_time = int((run.completed_at - run.started_at).total_seconds() * 1000)
|
||
|
||
# 发送 complete 事件
|
||
yield {
|
||
'type': 'complete',
|
||
'run_id': str(run.id),
|
||
'status': 'completed',
|
||
'outputs': run.outputs,
|
||
'elapsed_time': run.elapsed_time,
|
||
'total_tokens': total_tokens,
|
||
}
|
||
break
|
||
|
||
# 创建节点实例
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if not node_instance:
|
||
logger.warning(f'未知节点类型: {node_type}')
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
continue
|
||
|
||
# 执行节点(在线程池中运行同步代码)
|
||
start_time = time.time()
|
||
|
||
# LLM 节点:根据配置决定是否使用流式执行
|
||
node_data = node_config.get('data', {})
|
||
use_streaming = node_data.get('use_streaming', False) if node_type == 'llm' else False
|
||
if node_type == 'llm' and use_streaming and hasattr(node_instance, 'execute_stream'):
|
||
result = None
|
||
stream_gen = node_instance.execute_stream(context)
|
||
|
||
try:
|
||
while True:
|
||
# 在线程池中获取下一个chunk
|
||
chunk_event = await asyncio.get_event_loop().run_in_executor(
|
||
None, next, stream_gen
|
||
)
|
||
# 发送 LLM 流式内容事件
|
||
if chunk_event.content:
|
||
yield _merge_event_collaboration({
|
||
'type': 'llm_chunk',
|
||
'node_id': current_node_id,
|
||
'content': chunk_event.content,
|
||
'accumulated_content': chunk_event.accumulated_content,
|
||
}, node_map, current_node_id, node_type)
|
||
except StopIteration as e:
|
||
# 生成器结束,获取返回值(NodeResult)
|
||
result = e.value
|
||
|
||
if result is None:
|
||
result = NodeResult(success=False, error='LLM 节点执行异常')
|
||
else:
|
||
# 其他节点使用异步执行
|
||
result = await node_instance.execute_async(context)
|
||
|
||
elapsed = int((time.time() - start_time) * 1000)
|
||
target_node_id = _resolve_handoff_target(
|
||
current_node_id,
|
||
result,
|
||
edge_map,
|
||
parallel_edge_map,
|
||
node_map,
|
||
)
|
||
|
||
# 处理节点事件(如消息节点发送消息)
|
||
if result.events:
|
||
for event in result.events:
|
||
if event.get('type') == 'message':
|
||
# 消息事件直接发送为 answer 类型
|
||
answer_event = _merge_event_collaboration({
|
||
'type': 'answer',
|
||
'node_id': current_node_id,
|
||
'content': event.get('content', ''),
|
||
}, node_map, current_node_id, node_type, result)
|
||
if event.get('from_subflow'):
|
||
answer_event['from_subflow'] = True
|
||
answer_event['subflow_name'] = event.get('subflow_name')
|
||
yield answer_event
|
||
else:
|
||
# 其他事件保持 node_event 格式
|
||
yield _merge_event_collaboration({
|
||
'type': 'node_event',
|
||
'node_id': current_node_id,
|
||
'event': event,
|
||
}, node_map, current_node_id, node_type, result)
|
||
|
||
# 检查是否需要等待用户输入(对话流节点)
|
||
if result.waiting_for_input:
|
||
# 记录日志
|
||
log_entry = _make_execution_log_entry(
|
||
node_map,
|
||
current_node_id,
|
||
node_type,
|
||
status='waiting',
|
||
output=result.output,
|
||
elapsed_time=elapsed,
|
||
metadata=_node_result_metadata(result),
|
||
inputs=copy.deepcopy(node_inputs),
|
||
target_node_id=target_node_id,
|
||
)
|
||
logs.append(log_entry)
|
||
|
||
# 调试:打印等待时保存的变量
|
||
logger.info(f'等待用户输入,保存变量到 run.outputs,变量键: {list(context.variables.keys())}')
|
||
node_keys = [k for k in context.variables.keys() if k.startswith('_node_')]
|
||
logger.info(f'等待用户输入,_node_ 前缀的变量: {node_keys}')
|
||
|
||
# 更新运行记录为等待状态
|
||
run.status = 'waiting'
|
||
run.current_node_id = current_node_id
|
||
run.execution_log = logs.copy() # 使用 copy() 确保 SQLAlchemy 检测到变化
|
||
run.outputs = dict(context.variables) # 使用 dict() 创建新对象,确保 SQLAlchemy 检测到变化
|
||
run.waiting_config = dict(result.waiting_config) if result.waiting_config else {}
|
||
# 注意:save由调用方处理
|
||
|
||
# 发送等待输入事件
|
||
yield _merge_event_collaboration({
|
||
'type': 'waiting_input',
|
||
'node_id': current_node_id,
|
||
'node_type': node_type,
|
||
'config': result.waiting_config,
|
||
'waiting_config': result.waiting_config,
|
||
}, node_map, current_node_id, node_type, result)
|
||
|
||
# 暂停工作流执行,等待用户输入后续流
|
||
return
|
||
|
||
# 记录日志
|
||
log_entry = _make_execution_log_entry(
|
||
node_map,
|
||
current_node_id,
|
||
node_type,
|
||
status='completed' if result.success else 'failed',
|
||
output=result.output,
|
||
error=result.error,
|
||
elapsed_time=elapsed,
|
||
tokens_used=result.tokens_used,
|
||
metadata=_node_result_metadata(result),
|
||
inputs=copy.deepcopy(node_inputs),
|
||
target_node_id=target_node_id,
|
||
)
|
||
logs.append(log_entry)
|
||
|
||
total_tokens += result.tokens_used
|
||
total_steps += 1
|
||
|
||
# 更新运行记录
|
||
run.current_node_id = current_node_id
|
||
run.execution_log = logs.copy() # 使用 copy() 确保 SQLAlchemy 检测到变化
|
||
run.total_tokens = total_tokens
|
||
run.total_steps = total_steps
|
||
run.outputs = dict(context.variables) # 使用 dict() 创建新对象,确保 SQLAlchemy 检测到变化
|
||
# 注意:save由调用方处理
|
||
|
||
# 发送节点完成事件(包含输出数据)
|
||
yield _merge_event_collaboration(_apply_node_result_warnings({
|
||
'type': 'node_complete',
|
||
'node_id': current_node_id,
|
||
'node_type': node_type,
|
||
'status': 'success' if result.success else 'failed',
|
||
'elapsed_time': elapsed,
|
||
'tokens_used': result.tokens_used,
|
||
'error_message': result.error if not result.success else None,
|
||
'outputs': {
|
||
'output': result.output,
|
||
'output_variables': result.metadata.get('frontend_output_variables', result.output_variables),
|
||
},
|
||
'target_node_id': target_node_id,
|
||
}, result), node_map, current_node_id, node_type, result)
|
||
|
||
if not result.success:
|
||
raise ValueError(f'节点执行失败: {result.error}')
|
||
|
||
# 合并输出变量到上下文
|
||
context.variables.update(result.output_variables)
|
||
context.previous_output = result.output
|
||
|
||
# 保存节点输出到命名空间(支持 {{NodeID.key}} 格式引用)
|
||
context.variables[f'_node_{current_node_id}'] = {
|
||
'output': result.output,
|
||
**result.output_variables,
|
||
}
|
||
|
||
# 确定下一个节点
|
||
if node_type == 'parallel':
|
||
# 并行节点:执行所有分支(流式)
|
||
branch_ids = list(parallel_edge_map.get(current_node_id, {}).keys()) or edge_map.get(current_node_id, [])
|
||
branch_labels = _parallel_branch_labels(node_map, current_node_id)
|
||
yield _merge_event_collaboration({
|
||
'type': 'parallel_start',
|
||
'node_id': current_node_id,
|
||
'branches': branch_ids,
|
||
'branch_count': len(branch_ids),
|
||
'branch_labels': branch_labels,
|
||
'collaboration_mode': 'parallel',
|
||
}, node_map, current_node_id, node_type)
|
||
|
||
# 使用流式版本的并行分支执行(异步)
|
||
async for event in self._execute_parallel_branches_stream(
|
||
node_map, edge_map, parallel_edge_map, current_node_id, context, run, logs
|
||
):
|
||
if event.get('type') == '_parallel_result':
|
||
# 内部结果事件,不 yield 给前端
|
||
parallel_results = event['data']
|
||
context.set_variable('_parallel_results', parallel_results['results'])
|
||
total_tokens += parallel_results['total_tokens']
|
||
total_steps += parallel_results['total_steps']
|
||
logs.extend(parallel_results['logs'])
|
||
|
||
yield _merge_event_collaboration({
|
||
'type': 'parallel_complete',
|
||
'node_id': current_node_id,
|
||
'branch_results': parallel_results['results'],
|
||
'total_tokens': parallel_results['total_tokens'],
|
||
'branch_count': len(parallel_results['results']),
|
||
'branch_labels': branch_labels,
|
||
'collaboration_mode': 'parallel',
|
||
}, node_map, current_node_id, node_type)
|
||
|
||
# 找到合并节点继续执行
|
||
current_node_id = parallel_results.get('merge_node_id')
|
||
else:
|
||
# 其他事件(node_start, node_complete 等)直接 yield
|
||
yield event
|
||
elif node_type == 'loop':
|
||
# 循环节点:执行循环体
|
||
loop_node_id = current_node_id
|
||
loop_config = node_config.get('data', {})
|
||
|
||
yield {
|
||
'type': 'loop_start',
|
||
'node_id': loop_node_id,
|
||
'loop_mode': context.get_variable('_loop_mode'),
|
||
'total_iterations': context.get_variable('_loop_total', 0),
|
||
}
|
||
|
||
# 执行循环(异步)
|
||
loop_waiting = False
|
||
loop_error = None
|
||
async for loop_event in self._execute_loop(
|
||
node_map, edge_map, parallel_edge_map, loop_node_id,
|
||
loop_config, context, run, logs
|
||
):
|
||
yield loop_event
|
||
# 累加循环中的 token 和步骤
|
||
if loop_event.get('type') == 'loop_iteration_complete':
|
||
total_tokens += loop_event.get('tokens_used', 0)
|
||
total_steps += loop_event.get('steps', 0)
|
||
# 检查是否因等待用户输入而暂停
|
||
if loop_event.get('type') == 'waiting_input':
|
||
loop_waiting = True
|
||
# 检查是否有循环迭代错误
|
||
if loop_event.get('type') == 'loop_iteration_error':
|
||
loop_error = loop_event.get('error', '循环执行失败')
|
||
logger.error(f'[流式] 检测到循环迭代错误: {loop_error}')
|
||
|
||
# 如果循环因等待用户输入而暂停,停止主工作流执行
|
||
if loop_waiting:
|
||
return
|
||
|
||
# 如果循环中有错误,抛出异常
|
||
if loop_error:
|
||
raise ValueError(f'循环节点执行失败: {loop_error}')
|
||
|
||
# 获取循环结果
|
||
loop_results = context.get_variable('_loop_results', [])
|
||
output_var = loop_config.get('output_variable', 'loop_results')
|
||
context.set_variable(output_var, loop_results)
|
||
|
||
yield {
|
||
'type': 'loop_complete',
|
||
'node_id': loop_node_id,
|
||
'total_iterations': context.get_variable('_loop_index', 0),
|
||
'results_count': len(loop_results),
|
||
'loop_results': loop_results,
|
||
'output_variable': output_var,
|
||
}
|
||
|
||
# 单节点循环设计:循环完成后从循环节点的输出边继续
|
||
next_nodes = edge_map.get(loop_node_id, [])
|
||
# 过滤掉子节点(parentNode 为循环节点的节点)
|
||
child_node_ids = self._get_child_nodes(node_map, loop_node_id)
|
||
next_nodes = [n for n in next_nodes if n not in child_node_ids]
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
elif result.next_node_id:
|
||
# 检查是否是分支 ID(条件分支节点返回的是分支 ID,需要通过边映射找到实际目标节点)
|
||
logger.info(f'[流式] 节点 {current_node_id} 返回 next_node_id={result.next_node_id}')
|
||
logger.info(f'[流式] parallel_edge_map[{current_node_id}]={parallel_edge_map.get(current_node_id, {})}')
|
||
if result.next_node_id in parallel_edge_map.get(current_node_id, {}):
|
||
current_node_id = parallel_edge_map[current_node_id][result.next_node_id]
|
||
logger.info(f'[流式] 通过分支映射找到目标节点: {current_node_id}')
|
||
elif result.next_node_id in node_map:
|
||
# 是实际的节点 ID
|
||
current_node_id = result.next_node_id
|
||
logger.info(f'[流式] 直接使用节点 ID: {current_node_id}')
|
||
else:
|
||
# 尝试从边映射中查找
|
||
logger.warning(f'[流式] 分支 ID {result.next_node_id} 未找到对应边,回退到普通边查找')
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
else:
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
|
||
async def resume_workflow_stream(
|
||
self,
|
||
workflow: AIWorkflow,
|
||
run: AIWorkflowRun,
|
||
user_input: Any,
|
||
conversation_history: List[Dict[str, str]] = None,
|
||
):
|
||
"""
|
||
恢复等待中的工作流执行(异步生成器)
|
||
|
||
注意:此方法需要预先加载好workflow和run对象
|
||
|
||
Args:
|
||
workflow: 工作流对象
|
||
run: 工作流运行记录
|
||
user_input: 用户输入的值
|
||
conversation_history: 对话历史(用于 LLM 节点的上下文记忆)
|
||
|
||
Yields:
|
||
dict: 节点执行状态事件
|
||
"""
|
||
# 发送恢复事件
|
||
yield {
|
||
'type': 'resume',
|
||
'run_id': str(run.id),
|
||
'node_id': run.current_node_id,
|
||
}
|
||
|
||
try:
|
||
# 恢复执行(异步)
|
||
async for event in self._resume_workflow_stream(workflow, run, user_input, conversation_history):
|
||
yield event
|
||
|
||
# 发送完成事件(如果已完成)
|
||
if run.status == 'completed':
|
||
yield {
|
||
'type': 'complete',
|
||
'run_id': str(run.id),
|
||
'status': 'completed',
|
||
'outputs': run.outputs,
|
||
'elapsed_time': run.elapsed_time,
|
||
'total_tokens': run.total_tokens,
|
||
}
|
||
|
||
except Exception as e:
|
||
logger.exception(f'工作流恢复执行失败: {e}')
|
||
yield {
|
||
'type': 'error',
|
||
'run_id': str(run.id),
|
||
'node_id': run.current_node_id,
|
||
'message': str(e),
|
||
}
|
||
|
||
async def resume_workflow_stream_async(
|
||
self,
|
||
run_id: str,
|
||
user_input: Any,
|
||
conversation_history: List[Dict[str, str]] = None,
|
||
):
|
||
"""
|
||
异步恢复等待中的工作流执行
|
||
|
||
Args:
|
||
run_id: 运行记录 ID
|
||
user_input: 用户输入的值
|
||
conversation_history: 对话历史
|
||
|
||
Yields:
|
||
dict: 节点执行状态事件
|
||
"""
|
||
if not self._db:
|
||
yield {'type': 'error', 'message': '数据库会话未初始化'}
|
||
return
|
||
|
||
# 查询运行记录
|
||
result = await self._db.execute(
|
||
select(AIWorkflowRun).where(AIWorkflowRun.id == run_id)
|
||
)
|
||
run = result.scalar_one_or_none()
|
||
if not run:
|
||
yield {'type': 'error', 'message': f'运行记录不存在: {run_id}'}
|
||
return
|
||
|
||
if run.status != 'waiting':
|
||
yield {'type': 'error', 'message': f'工作流不在等待状态: {run.status}'}
|
||
return
|
||
|
||
# 查询工作流
|
||
workflow = await self.get_workflow(run.workflow_id)
|
||
if not workflow:
|
||
yield {'type': 'error', 'message': '工作流不存在'}
|
||
return
|
||
|
||
try:
|
||
# 使用异步生成器恢复执行
|
||
async for event in self.resume_workflow_stream(workflow, run, user_input, conversation_history):
|
||
yield event
|
||
|
||
# 更新运行记录
|
||
if run.status == 'completed':
|
||
# 更新工作流统计
|
||
workflow.run_count = (workflow.run_count or 0) + 1
|
||
workflow.success_count = (workflow.success_count or 0) + 1
|
||
|
||
await self._db.commit()
|
||
|
||
except Exception as e:
|
||
logger.exception(f'工作流恢复执行失败: {e}')
|
||
run.status = 'failed'
|
||
run.error_message = str(e)
|
||
run.completed_at = datetime.now()
|
||
run.elapsed_time = int((run.completed_at - run.started_at).total_seconds() * 1000)
|
||
|
||
await self._db.commit()
|
||
|
||
# 发送错误事件给前端
|
||
yield {
|
||
'type': 'error',
|
||
'message': str(e),
|
||
'error_message': str(e),
|
||
'run_id': run.id,
|
||
'workflow_id': workflow.id,
|
||
}
|
||
|
||
async def _resume_workflow_stream(
|
||
self,
|
||
workflow: AIWorkflow,
|
||
run: AIWorkflowRun,
|
||
user_input: Any,
|
||
conversation_history: List[Dict[str, str]] = None,
|
||
):
|
||
"""
|
||
恢复工作流执行的内部实现(异步版本)
|
||
|
||
Args:
|
||
workflow: 工作流对象
|
||
run: 工作流运行记录
|
||
user_input: 用户输入
|
||
conversation_history: 对话历史(用于 LLM 节点的上下文记忆)
|
||
"""
|
||
import asyncio
|
||
definition = workflow.definition
|
||
nodes = definition.get('nodes', [])
|
||
edges = definition.get('edges', [])
|
||
|
||
# 构建节点映射和边映射
|
||
node_map = {node['id']: node for node in nodes}
|
||
edge_map = {}
|
||
parallel_edge_map = {}
|
||
for edge in edges:
|
||
source = edge.get('source', '')
|
||
target = edge.get('target', '')
|
||
source_handle = edge.get('sourceHandle', '')
|
||
if source not in edge_map:
|
||
edge_map[source] = []
|
||
edge_map[source].append(target)
|
||
if source_handle:
|
||
if source not in parallel_edge_map:
|
||
parallel_edge_map[source] = {}
|
||
parallel_edge_map[source][source_handle] = target
|
||
|
||
# 恢复上下文
|
||
logger.info(f'[WorkflowService] 恢复工作流,接收到的 user_input 类型: {type(user_input)}, 值: {user_input}')
|
||
|
||
context = NodeContext(
|
||
workflow_run_id=str(run.id),
|
||
variables=run.outputs.copy() if run.outputs else {},
|
||
user_input=str(user_input) if user_input is not None else '',
|
||
user_id=str(run.user_id) if run.user_id else '',
|
||
conversation_history=conversation_history or [],
|
||
db_session=self._db,
|
||
)
|
||
|
||
logger.info(f'[WorkflowService] context.user_input 设置为: {context.user_input}')
|
||
|
||
# 调试:打印恢复时加载的变量
|
||
logger.info(f'恢复工作流,从 run.outputs 加载变量,变量键: {list(context.variables.keys())}')
|
||
node_keys = [k for k in context.variables.keys() if k.startswith('_node_')]
|
||
logger.info(f'恢复工作流,_node_ 前缀的变量: {node_keys}')
|
||
# 检查循环相关变量
|
||
loop_items = context.variables.get('_loop_items')
|
||
logger.info(f'恢复工作流,_loop_items 长度={len(loop_items) if loop_items else 0}, _loop_index={context.variables.get("_loop_index")}, _loop_mode={context.variables.get("_loop_mode")}')
|
||
|
||
# 检查是否有旧的 __user_input__
|
||
old_user_input = context.variables.get('__user_input__')
|
||
logger.info(f'[WorkflowService] variables 中的旧 __user_input__: {old_user_input}')
|
||
|
||
# 设置用户输入到上下文(对话流节点会读取这个值)
|
||
context.variables['__user_input__'] = user_input
|
||
logger.info(f'[WorkflowService] 更新 variables["__user_input__"] 为: {user_input}')
|
||
|
||
logs = run.execution_log or []
|
||
total_tokens = run.total_tokens
|
||
total_steps = run.total_steps
|
||
current_node_id = run.current_node_id
|
||
|
||
# 检查是否有循环状态需要恢复
|
||
loop_state = run.waiting_config.get('_loop_state') if run.waiting_config else None
|
||
# 检查是否有子流程状态需要恢复
|
||
subflow_state = run.waiting_config.get('_subflow_state') if run.waiting_config else None
|
||
logger.info(f'恢复工作流,loop_state={loop_state}, subflow_state={bool(subflow_state)}, waiting_config={run.waiting_config}')
|
||
|
||
# 如果有子流程状态,将其添加到上下文的 metadata 中
|
||
if subflow_state:
|
||
context.metadata['_subflow_state'] = subflow_state
|
||
|
||
# 更新状态为运行中
|
||
run.status = 'running'
|
||
run.waiting_config = {}
|
||
|
||
while current_node_id:
|
||
node_config = node_map.get(current_node_id)
|
||
if not node_config:
|
||
break
|
||
|
||
node_type = node_config.get('type', '')
|
||
node_label = node_config.get('data', {}).get('label', current_node_id)
|
||
|
||
# 获取节点输入
|
||
node_inputs = {
|
||
'variables': dict(context.variables),
|
||
'previous_output': context.previous_output,
|
||
}
|
||
|
||
# 发送节点开始事件
|
||
node_start_event = _merge_event_collaboration({
|
||
'type': 'node_start',
|
||
'node_id': current_node_id,
|
||
'node_type': node_type,
|
||
'node_label': node_label,
|
||
'timestamp': datetime.now().isoformat(),
|
||
'inputs': node_inputs,
|
||
}, node_map, current_node_id, node_type)
|
||
# 如果在循环中,添加迭代信息
|
||
if loop_state:
|
||
node_start_event['loop_iteration'] = loop_state.get('iteration', 0)
|
||
yield node_start_event
|
||
|
||
# 结束节点
|
||
if node_type == 'end':
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if node_instance:
|
||
result = node_instance.execute(context)
|
||
logs.append(_make_execution_log_entry(
|
||
node_map,
|
||
current_node_id,
|
||
node_type,
|
||
status='completed',
|
||
output=result.output,
|
||
inputs=copy.deepcopy(node_inputs),
|
||
))
|
||
yield _merge_event_collaboration({
|
||
'type': 'node_complete',
|
||
'node_id': current_node_id,
|
||
'node_type': node_type,
|
||
'status': 'success',
|
||
'elapsed_time': 0,
|
||
'outputs': {'output': result.output},
|
||
}, node_map, current_node_id, node_type, result)
|
||
|
||
# 更新运行记录,保存结束节点的输出
|
||
run.status = 'completed'
|
||
run.completed_at = datetime.now()
|
||
run.elapsed_time = int((run.completed_at - run.started_at).total_seconds() * 1000)
|
||
run.execution_log = logs.copy()
|
||
# 只有当结束节点有输出时才包含 output 字段
|
||
if result.output is not None:
|
||
run.outputs = {
|
||
**context.variables,
|
||
'output': result.output, # 结束节点的输出内容
|
||
}
|
||
else:
|
||
run.outputs = dict(context.variables)
|
||
# 注意:save由调用方处理
|
||
break
|
||
|
||
# 创建节点实例并执行(在线程池中运行同步代码)
|
||
node_instance = NodeRegistry.create_instance(node_type, node_config.get('data', {}))
|
||
if not node_instance:
|
||
raise ValueError(f'未知节点类型: {node_type}')
|
||
|
||
start_time = time.time()
|
||
result = await node_instance.execute_async(context)
|
||
elapsed = int((time.time() - start_time) * 1000)
|
||
|
||
# 处理节点事件
|
||
if result.events:
|
||
for event in result.events:
|
||
if event.get('type') == 'message':
|
||
# 消息事件直接发送为 answer 类型
|
||
answer_event = _merge_event_collaboration({
|
||
'type': 'answer',
|
||
'node_id': current_node_id,
|
||
'content': event.get('content', ''),
|
||
}, node_map, current_node_id, node_type, result)
|
||
if event.get('from_subflow'):
|
||
answer_event['from_subflow'] = True
|
||
answer_event['subflow_name'] = event.get('subflow_name')
|
||
yield answer_event
|
||
else:
|
||
# 其他事件保持 node_event 格式
|
||
yield _merge_event_collaboration({
|
||
'type': 'node_event',
|
||
'node_id': current_node_id,
|
||
'event': event,
|
||
}, node_map, current_node_id, node_type, result)
|
||
|
||
# 检查是否需要等待用户输入
|
||
if result.waiting_for_input:
|
||
log_entry = _make_execution_log_entry(
|
||
node_map,
|
||
current_node_id,
|
||
node_type,
|
||
status='waiting',
|
||
output=result.output,
|
||
elapsed_time=elapsed,
|
||
metadata=_node_result_metadata(result),
|
||
target_node_id=target_node_id,
|
||
)
|
||
logs.append(log_entry)
|
||
|
||
run.status = 'waiting'
|
||
run.current_node_id = current_node_id
|
||
run.execution_log = logs.copy() # 使用 copy() 确保 SQLAlchemy 检测到变化
|
||
run.outputs = dict(context.variables) # 使用 dict() 创建新对象,确保 SQLAlchemy 检测到变化
|
||
run.waiting_config = dict(result.waiting_config) if result.waiting_config else {}
|
||
# 如果在循环中,保存循环状态以便恢复
|
||
if loop_state:
|
||
run.waiting_config['_loop_state'] = loop_state
|
||
# 注意:save由调用方处理
|
||
|
||
waiting_event = _merge_event_collaboration({
|
||
'type': 'waiting_input',
|
||
'node_id': current_node_id,
|
||
'node_type': node_type,
|
||
'config': result.waiting_config,
|
||
'waiting_config': result.waiting_config,
|
||
}, node_map, current_node_id, node_type, result)
|
||
# 如果在循环中,添加迭代信息
|
||
if loop_state:
|
||
waiting_event['loop_iteration'] = loop_state.get('iteration', 0)
|
||
yield waiting_event
|
||
return
|
||
|
||
# 记录日志
|
||
log_entry = _make_execution_log_entry(
|
||
node_map,
|
||
current_node_id,
|
||
node_type,
|
||
status='completed' if result.success else 'failed',
|
||
output=result.output,
|
||
error=result.error,
|
||
elapsed_time=elapsed,
|
||
tokens_used=result.tokens_used,
|
||
metadata=_node_result_metadata(result),
|
||
target_node_id=target_node_id,
|
||
)
|
||
logs.append(log_entry)
|
||
|
||
total_tokens += result.tokens_used
|
||
total_steps += 1
|
||
|
||
# 清除用户输入(只在第一个节点使用)
|
||
context.variables.pop('__user_input__', None)
|
||
|
||
# 更新运行记录
|
||
run.current_node_id = current_node_id
|
||
run.execution_log = logs.copy() # 使用 copy() 确保 SQLAlchemy 检测到变化
|
||
run.total_tokens = total_tokens
|
||
run.total_steps = total_steps
|
||
run.outputs = dict(context.variables) # 使用 dict() 创建新对象,确保 SQLAlchemy 检测到变化
|
||
# 注意:save由调用方处理
|
||
|
||
# 发送节点完成事件
|
||
node_complete_event = _merge_event_collaboration(_apply_node_result_warnings({
|
||
'type': 'node_complete',
|
||
'node_id': current_node_id,
|
||
'node_type': node_type,
|
||
'status': 'success' if result.success else 'failed',
|
||
'elapsed_time': elapsed,
|
||
'tokens_used': result.tokens_used,
|
||
'error_message': result.error if not result.success else None,
|
||
'outputs': {
|
||
'output': result.output,
|
||
'output_variables': result.metadata.get('frontend_output_variables', result.output_variables),
|
||
},
|
||
'target_node_id': target_node_id,
|
||
}, result), node_map, current_node_id, node_type, result)
|
||
# 如果在循环中,添加迭代信息
|
||
if loop_state:
|
||
node_complete_event['loop_iteration'] = loop_state.get('iteration', 0)
|
||
yield node_complete_event
|
||
|
||
if not result.success:
|
||
raise ValueError(f'节点执行失败: {result.error}')
|
||
|
||
# 合并输出变量到上下文
|
||
context.variables.update(result.output_variables)
|
||
context.previous_output = result.output
|
||
|
||
# 保存节点输出到命名空间
|
||
context.variables[f'_node_{current_node_id}'] = {
|
||
'output': result.output,
|
||
**result.output_variables,
|
||
}
|
||
|
||
# 确定下一个节点
|
||
# 检查当前节点是否是循环节点
|
||
if node_type == 'loop':
|
||
# 循环节点:执行循环体
|
||
loop_node_id = current_node_id
|
||
loop_config = node_config.get('data', {})
|
||
|
||
yield {
|
||
'type': 'loop_start',
|
||
'node_id': loop_node_id,
|
||
'loop_mode': context.get_variable('_loop_mode'),
|
||
'total_iterations': context.get_variable('_loop_total', 0),
|
||
}
|
||
|
||
# 执行循环(异步)
|
||
loop_waiting = False
|
||
async for loop_event in self._execute_loop(
|
||
node_map, edge_map, parallel_edge_map, loop_node_id,
|
||
loop_config, context, run, logs
|
||
):
|
||
yield loop_event
|
||
# 累加循环中的 token 和步骤
|
||
if loop_event.get('type') == 'loop_iteration_complete':
|
||
total_tokens += loop_event.get('tokens_used', 0)
|
||
total_steps += loop_event.get('steps', 0)
|
||
# 检查是否因等待用户输入而暂停
|
||
if loop_event.get('type') == 'waiting_input':
|
||
loop_waiting = True
|
||
|
||
# 如果循环因等待用户输入而暂停,停止主工作流执行
|
||
if loop_waiting:
|
||
return
|
||
|
||
# 获取循环结果
|
||
loop_results = context.get_variable('_loop_results', [])
|
||
output_var = loop_config.get('output_variable', 'loop_results')
|
||
context.set_variable(output_var, loop_results)
|
||
|
||
yield {
|
||
'type': 'loop_complete',
|
||
'node_id': loop_node_id,
|
||
'total_iterations': context.get_variable('_loop_index', 0),
|
||
'results_count': len(loop_results),
|
||
'loop_results': loop_results,
|
||
'output_variable': output_var,
|
||
}
|
||
|
||
# 单节点循环设计:循环完成后从循环节点的输出边继续
|
||
next_nodes = edge_map.get(loop_node_id, [])
|
||
# 过滤掉子节点(parentNode 为循环节点的节点)
|
||
child_node_ids = self._get_child_nodes(node_map, loop_node_id)
|
||
next_nodes = [n for n in next_nodes if n not in child_node_ids]
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
# 检查是否在循环中恢复执行
|
||
elif loop_state:
|
||
loop_node_id = loop_state.get('loop_node_id')
|
||
loop_node_config = node_map.get(loop_node_id, {})
|
||
loop_config = loop_node_config.get('data', {})
|
||
|
||
# 获取循环体子节点
|
||
child_node_ids = self._get_child_nodes(node_map, loop_node_id)
|
||
if not child_node_ids:
|
||
child_node_ids = loop_config.get('childNodes', [])
|
||
|
||
# 构建子节点的边映射
|
||
child_edge_map = {}
|
||
for source_id in child_node_ids:
|
||
targets = edge_map.get(source_id, [])
|
||
child_targets = [t for t in targets if t in child_node_ids]
|
||
if child_targets:
|
||
child_edge_map[source_id] = child_targets
|
||
|
||
# 找循环体内的下一个节点
|
||
if result.next_node_id and result.next_node_id in child_node_ids:
|
||
current_node_id = result.next_node_id
|
||
else:
|
||
next_nodes = child_edge_map.get(current_node_id, [])
|
||
if next_nodes:
|
||
current_node_id = next_nodes[0]
|
||
else:
|
||
# 循环体内没有更多节点,完成当前迭代并继续循环
|
||
logger.info(f'[Resume] 循环体内节点执行完毕,完成当前迭代')
|
||
|
||
# 收集当前迭代结果
|
||
loop_results = context.get_variable('_loop_results', [])
|
||
loop_results.append(result.output)
|
||
context.set_variable('_loop_results', loop_results)
|
||
|
||
# 更新循环索引(进入下一次迭代)
|
||
current_loop_index = context.get_variable('_loop_index', 0)
|
||
context.set_variable('_loop_index', current_loop_index + 1)
|
||
logger.info(f'[Resume] 更新 _loop_index: {current_loop_index} -> {current_loop_index + 1}')
|
||
|
||
# 发送迭代完成事件
|
||
loop_items = context.get_variable('_loop_items', [])
|
||
yield {
|
||
'type': 'loop_iteration_complete',
|
||
'node_id': loop_node_id,
|
||
'iteration': current_loop_index,
|
||
'total': len(loop_items) if loop_items else None,
|
||
'output': result.output,
|
||
}
|
||
|
||
# 继续循环执行
|
||
loop_items = context.get_variable('_loop_items', [])
|
||
loop_index = context.get_variable('_loop_index', 0)
|
||
logger.info(f'[Resume] 继续循环执行, _loop_items 长度={len(loop_items) if loop_items else 0}, _loop_index={loop_index}')
|
||
logger.info(f'[Resume] context.variables 中的循环变量: _loop_items={bool(context.variables.get("_loop_items"))}, _loop_mode={context.variables.get("_loop_mode")}')
|
||
loop_waiting = False
|
||
async for loop_event in self._execute_loop(
|
||
node_map, edge_map, parallel_edge_map, loop_node_id,
|
||
loop_config, context, run, logs
|
||
):
|
||
yield loop_event
|
||
if loop_event.get('type') == 'waiting_input':
|
||
loop_waiting = True
|
||
|
||
# 如果循环因等待用户输入而暂停,停止执行
|
||
if loop_waiting:
|
||
return
|
||
|
||
# 循环执行完成,找循环节点之后的节点
|
||
next_nodes = edge_map.get(loop_node_id, [])
|
||
next_nodes = [n for n in next_nodes if n not in child_node_ids]
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
loop_state = None # 清除循环状态
|
||
elif result.next_node_id:
|
||
# 检查是否是分支 ID(条件分支节点返回的是分支 ID,需要通过边映射找到实际目标节点)
|
||
if result.next_node_id in parallel_edge_map.get(current_node_id, {}):
|
||
current_node_id = parallel_edge_map[current_node_id][result.next_node_id]
|
||
elif result.next_node_id in node_map:
|
||
# 是实际的节点 ID
|
||
current_node_id = result.next_node_id
|
||
else:
|
||
# 尝试从边映射中查找
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
else:
|
||
next_nodes = edge_map.get(current_node_id, [])
|
||
current_node_id = next_nodes[0] if next_nodes else None
|
||
|
||
@staticmethod
|
||
def get_node_schemas() -> List[Dict]:
|
||
"""获取所有节点 Schema"""
|
||
return NodeRegistry.get_all_schemas()
|
||
|
||
@staticmethod
|
||
def get_node_schemas_by_category() -> Dict[str, List[Dict]]:
|
||
"""按分类获取节点 Schema"""
|
||
return NodeRegistry.get_schemas_by_category()
|