""" 节点注册中心 """ 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}')