Restore AI workflow design nodes
This commit is contained in:
@@ -12,53 +12,53 @@ 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
|
||||
"""
|
||||
@@ -66,24 +66,24 @@ class NodeRegistry:
|
||||
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
|
||||
"""
|
||||
@@ -94,12 +94,12 @@ class NodeRegistry:
|
||||
result[category] = []
|
||||
result[category].append(node.get_schema())
|
||||
return result
|
||||
|
||||
|
||||
@classmethod
|
||||
def get_all_types(cls) -> List[str]:
|
||||
"""
|
||||
获取所有已注册的节点类型
|
||||
|
||||
|
||||
Returns:
|
||||
节点类型列表
|
||||
"""
|
||||
@@ -130,6 +130,23 @@ def _load_builtin_nodes():
|
||||
'text_to_sql_node',
|
||||
'snowflake_cortex_node',
|
||||
'knowledge_retrieval_node',
|
||||
'form_basic_info_node',
|
||||
'form_database_design_node',
|
||||
'form_database_create_node',
|
||||
'form_ui_design_node',
|
||||
'form_list_design_node',
|
||||
'form_create_node',
|
||||
'form_publish_node',
|
||||
'form_data_node',
|
||||
'app_create_node',
|
||||
'app_design_node',
|
||||
'app_settings_node',
|
||||
'app_update_node',
|
||||
'dashboard_basic_info_node',
|
||||
'dashboard_design_node',
|
||||
'dashboard_create_node',
|
||||
'dashboard_publish_node',
|
||||
'system_summary_node',
|
||||
]
|
||||
|
||||
from importlib import import_module
|
||||
|
||||
Reference in New Issue
Block a user