Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,394 @@
|
||||
"""
|
||||
节点基类
|
||||
"""
|
||||
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
|
||||
Reference in New Issue
Block a user