Files
ai-agent-admin/backend-fastapi/ai_platform/nodes/base.py
T
2026-06-08 18:14:59 +08:00

395 lines
13 KiB
Python

"""
节点基类
"""
import logging
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
@dataclass
class NodeContext:
"""
节点执行上下文
包含节点执行所需的所有信息
"""
# 工作流运行实例
workflow_run_id: str = ''
# 变量存储(所有节点共享)
variables: Dict[str, Any] = field(default_factory=dict)
# 用户输入
user_input: str = ''
# 当前用户
user_id: str = ''
# 对话历史(用于 LLM 节点)
conversation_history: List[Dict[str, str]] = field(default_factory=list)
# 节点配置
node_config: Dict[str, Any] = field(default_factory=dict)
# 上一个节点的输出
previous_output: Any = None
# 元数据
metadata: Dict[str, Any] = field(default_factory=dict)
# 数据库会话(用于需要数据库访问的节点)
db_session: Any = None
def get_variable(self, name: str, default: Any = None) -> Any:
"""获取变量"""
return self.variables.get(name, default)
def set_variable(self, name: str, value: Any) -> None:
"""设置变量"""
self.variables[name] = value
def resolve_template(self, template: str) -> str:
"""
解析模板中的变量引用
支持格式:
- {{variable_name}} - 直接变量引用
- {{variable_name[0]}} - 数组索引访问
- {{variable_name[-1]}} - 负数索引(最后一个)
- {{NodeID.key}} - 节点输出引用(如果存在)
- {{NodeID.key}}.property - 访问解析结果的嵌套属性
- {{NodeID.key[0].property}} - 数组索引 + 属性访问
"""
import re
import json
result = template
# 匹配 {{...}} 格式的变量引用,以及可选的后续属性访问 .property1.property2...
pattern = r'\{\{([^}]+)\}\}((?:\.[a-zA-Z_][a-zA-Z0-9_]*)*)'
def get_nested_value(obj, path: str):
"""从对象中获取嵌套属性值,支持数组索引"""
if not path:
return obj
# 移除开头的点
if path.startswith('.'):
path = path[1:]
current = obj
# 使用正则分割路径,支持 .property 和 [index] 格式
# 例如: "items[0].name" -> ["items", "[0]", "name"]
parts = re.split(r'(?=\[)|\.', path)
for part in parts:
if not part:
continue
# 如果是字符串,尝试解析为 JSON
if isinstance(current, str):
try:
current = json.loads(current)
except (json.JSONDecodeError, TypeError):
return None
# 检查是否是数组索引 [n]
index_match = re.match(r'\[(-?\d+)\]', part)
if index_match:
index = int(index_match.group(1))
if isinstance(current, (list, tuple)):
try:
current = current[index]
except IndexError:
return None
else:
return None
# 从字典中获取属性
elif isinstance(current, dict):
if part in current:
current = current[part]
else:
return None
else:
return None
return current
def replace_var(match):
var_ref = match.group(1).strip()
extra_path = match.group(2) or '' # 额外的属性路径,如 .customer_name
value = None
index_path = '' # 变量名后的索引/属性路径
# 检查是否有数组索引 [n],如果有则分割
# 例如: "loop_results[0]" -> var_name="loop_results", index_path="[0]"
# 例如: "llm-123.llm_response" -> var_name="llm-123.llm_response", index_path=""
bracket_pos = var_ref.find('[')
if bracket_pos > 0:
var_name = var_ref[:bracket_pos]
index_path = var_ref[bracket_pos:]
else:
var_name = var_ref
# 尝试解析 NodeID.key.subkey... 格式(支持多层属性访问)
if '.' in var_name:
parts = var_name.split('.')
node_id = parts[0]
remaining_path = '.'.join(parts[1:]) # 剩余路径,如 "item.module_name"
# 先尝试从节点输出命名空间获取
node_outputs = self.variables.get(f'_node_{node_id}')
if isinstance(node_outputs, dict):
# 尝试获取第一层 key
first_key = parts[1] if len(parts) > 1 else None
if first_key and first_key in node_outputs:
value = node_outputs[first_key]
# 如果还有更多层级,继续递归获取
if len(parts) > 2:
nested_path = '.'.join(parts[2:])
nested_value = get_nested_value(value, nested_path)
if nested_value is not None:
value = nested_value
# 回退:尝试直接从变量获取(node_id 作为变量名)
if value is None and node_id in self.variables:
node_data = self.variables[node_id]
if isinstance(node_data, dict):
first_key = parts[1] if len(parts) > 1 else None
if first_key and first_key in node_data:
value = node_data[first_key]
if len(parts) > 2:
nested_path = '.'.join(parts[2:])
nested_value = get_nested_value(value, nested_path)
if nested_value is not None:
value = nested_value
# 再回退:直接从顶层变量获取完整路径
if value is None and remaining_path in self.variables:
value = self.variables[remaining_path]
else:
# 直接变量引用
if var_name in self.variables:
value = self.variables[var_name]
# 如果找到了值,处理索引路径和额外的属性路径
if value is not None:
# 合并索引路径和额外路径
full_path = index_path + extra_path
if full_path:
nested_value = get_nested_value(value, full_path)
if nested_value is not None:
return str(nested_value)
# 嵌套属性未找到,返回原始值
return str(value)
return str(value)
# 未找到变量,保持原样
return match.group(0)
result = re.sub(pattern, replace_var, result)
return result
@dataclass
class NodeResult:
"""
节点执行结果
"""
# 是否成功
success: bool = True
# 输出数据
output: Any = None
# 输出变量(会合并到上下文的 variables 中)
output_variables: Dict[str, Any] = field(default_factory=dict)
# 错误信息
error: str = ''
# 下一个节点 ID(用于条件分支)
next_node_id: str = ''
# Token 使用(LLM 节点)
tokens_used: int = 0
# 耗时(毫秒)
elapsed_time: int = 0
# 元数据
metadata: Dict[str, Any] = field(default_factory=dict)
# ========== 对话流相关 ==========
# 是否等待用户输入(对话流模式)
waiting_for_input: bool = False
# 等待配置(描述需要什么类型的输入)
waiting_config: Dict[str, Any] = field(default_factory=dict)
# 事件列表(如发送消息)
events: List[Dict[str, Any]] = field(default_factory=list)
# ========== 设计预览相关 ==========
# 设计预览数据(用于工作流中显示设计结果并允许编辑)
preview: Optional[Dict[str, Any]] = None
class BaseNode(ABC):
"""
节点基类
所有节点必须继承此类并实现 execute 方法
"""
# 节点类型标识(必须唯一)
node_type: str = ''
# 节点显示名称
node_name: str = ''
# 节点分类
node_category: str = 'basic' # basic, llm, logic, data, tool, knowledge
# 节点图标
node_icon: str = ''
# 节点描述
node_description: str = ''
# 输入参数定义
inputs: List[Dict[str, Any]] = []
# 输出参数定义
outputs: List[Dict[str, Any]] = []
# 是否支持多个输出分支
supports_branches: bool = False
def __init__(self, config: Dict[str, Any] = None):
"""
初始化节点
Args:
config: 节点配置
"""
self.config = config or {}
@abstractmethod
def execute(self, context: NodeContext) -> NodeResult:
"""
执行节点
Args:
context: 执行上下文
Returns:
NodeResult
"""
pass
async def execute_async(self, context: NodeContext) -> NodeResult:
"""
异步执行节点(默认调用同步方法)
Args:
context: 执行上下文
Returns:
NodeResult
"""
return self.execute(context)
def validate_config(self) -> tuple:
"""
验证节点配置
Returns:
(is_valid, error_message)
"""
return True, ''
@classmethod
def get_schema(cls) -> Dict[str, Any]:
"""
获取节点 Schema(供前端渲染)
Returns:
节点 Schema
"""
return {
'type': cls.node_type,
'name': cls.node_name,
'category': cls.node_category,
'icon': cls.node_icon,
'description': cls.node_description,
'inputs': cls.inputs,
'outputs': cls.outputs,
'supports_branches': cls.supports_branches,
}
@classmethod
def get_config_schema(cls) -> Dict[str, Any]:
"""
获取节点配置 Schema(供前端表单渲染)
Returns:
配置 Schema
"""
return {}
def resolve_require_confirmation(self, context: NodeContext, default: bool = True) -> bool:
"""
解析 require_confirmation 配置
支持三种模式:
1. 布尔值:直接使用 True/False
2. 字符串 'always'/'never':始终确认/从不确认
3. 变量引用:{{variable_name}} 格式,解析变量值作为布尔值
Args:
context: 节点执行上下文
default: 默认值
Returns:
是否需要确认
"""
value = self.config.get('require_confirmation', default)
# 布尔值直接返回
if isinstance(value, bool):
return value
# 字符串处理
if isinstance(value, str):
value_lower = value.lower().strip()
# 固定模式
if value_lower in ('always', 'true', '1', 'yes'):
return True
if value_lower in ('never', 'false', '0', 'no'):
return False
# 变量引用模式:{{variable_name}}
if '{{' in value and '}}' in value:
resolved = context.resolve_template(value)
# 解析后的值转换为布尔值
if isinstance(resolved, bool):
return resolved
if isinstance(resolved, str):
resolved_lower = resolved.lower().strip()
if resolved_lower in ('true', '1', 'yes'):
return True
if resolved_lower in ('false', '0', 'no'):
return False
# 非空字符串视为 True
return bool(resolved and resolved != value)
# 其他情况返回默认值
return default