Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,155 @@
|
||||
"""
|
||||
节点注册中心
|
||||
"""
|
||||
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}')
|
||||
Reference in New Issue
Block a user