156 lines
3.7 KiB
Python
156 lines
3.7 KiB
Python
"""
|
|
节点注册中心
|
|
"""
|
|
import logging
|
|
from typing import Dict, List, Optional, Type
|
|
|
|
from .base import BaseNode
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class NodeRegistry:
|
|
"""
|
|
节点注册中心
|
|
|
|
管理所有工作流节点的注册和获取
|
|
"""
|
|
|
|
_nodes: Dict[str, Type[BaseNode]] = {}
|
|
|
|
@classmethod
|
|
def register(cls, node_class: Type[BaseNode]) -> Type[BaseNode]:
|
|
"""
|
|
注册节点(可作为装饰器使用)
|
|
|
|
Args:
|
|
node_class: 节点类
|
|
|
|
Returns:
|
|
节点类
|
|
"""
|
|
node_type = node_class.node_type
|
|
if not node_type:
|
|
raise ValueError(f'Node {node_class.__name__} must have a node_type')
|
|
|
|
cls._nodes[node_type] = node_class
|
|
logger.info(f'Registered AI workflow node: {node_type}')
|
|
return node_class
|
|
|
|
@classmethod
|
|
def get(cls, node_type: str) -> Optional[Type[BaseNode]]:
|
|
"""
|
|
获取节点类
|
|
|
|
Args:
|
|
node_type: 节点类型
|
|
|
|
Returns:
|
|
节点类或 None
|
|
"""
|
|
return cls._nodes.get(node_type)
|
|
|
|
@classmethod
|
|
def create_instance(cls, node_type: str, config: Dict = None) -> Optional[BaseNode]:
|
|
"""
|
|
创建节点实例
|
|
|
|
Args:
|
|
node_type: 节点类型
|
|
config: 节点配置
|
|
|
|
Returns:
|
|
节点实例或 None
|
|
"""
|
|
node_class = cls.get(node_type)
|
|
if not node_class:
|
|
logger.warning(f'Unknown node type: {node_type}')
|
|
return None
|
|
|
|
return node_class(config=config)
|
|
|
|
@classmethod
|
|
def get_all_schemas(cls) -> List[Dict]:
|
|
"""
|
|
获取所有节点 Schema
|
|
|
|
Returns:
|
|
节点 Schema 列表
|
|
"""
|
|
return [node.get_schema() for node in cls._nodes.values()]
|
|
|
|
@classmethod
|
|
def get_schemas_by_category(cls) -> Dict[str, List[Dict]]:
|
|
"""
|
|
按分类获取节点 Schema
|
|
|
|
Returns:
|
|
按分类分组的节点 Schema
|
|
"""
|
|
result = {}
|
|
for node in cls._nodes.values():
|
|
category = node.node_category
|
|
if category not in result:
|
|
result[category] = []
|
|
result[category].append(node.get_schema())
|
|
return result
|
|
|
|
@classmethod
|
|
def get_all_types(cls) -> List[str]:
|
|
"""
|
|
获取所有已注册的节点类型
|
|
|
|
Returns:
|
|
节点类型列表
|
|
"""
|
|
return list(cls._nodes.keys())
|
|
|
|
|
|
# 自动加载所有内置节点
|
|
def _load_builtin_nodes():
|
|
"""加载所有内置节点"""
|
|
required_modules = [
|
|
'start_node',
|
|
'end_node',
|
|
'condition_node',
|
|
'template_node',
|
|
'parallel_node',
|
|
'merge_node',
|
|
]
|
|
optional_modules = [
|
|
'llm_node',
|
|
'code_node',
|
|
'http_node',
|
|
'variable_node',
|
|
'database_node',
|
|
'dialog_nodes',
|
|
'intent_node',
|
|
'loop_node',
|
|
'subflow_node',
|
|
'text_to_sql_node',
|
|
'snowflake_cortex_node',
|
|
'knowledge_retrieval_node',
|
|
]
|
|
|
|
from importlib import import_module
|
|
|
|
for module_name in required_modules:
|
|
import_module(f'{__package__}.builtin.{module_name}')
|
|
|
|
for module_name in optional_modules:
|
|
try:
|
|
import_module(f'{__package__}.builtin.{module_name}')
|
|
except ImportError as e:
|
|
logger.warning(
|
|
'Skipped optional AI workflow node module %s because dependencies are unavailable: %s',
|
|
module_name,
|
|
e,
|
|
)
|
|
|
|
|
|
# 延迟加载
|
|
try:
|
|
_load_builtin_nodes()
|
|
except ImportError as e:
|
|
logger.warning(f'Failed to load some builtin nodes: {e}')
|