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

156 lines
3.8 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}')