""" 核心注册机制模块 提供统一的注册器模式,支持动态扩展: - NodeRegistry: 节点注册器 - ToolRegistry: 工具注册器 - WorkflowRegistry: 工作流注册器 - ProviderRegistry: LLM Provider 注册器 """ from typing import Any, Callable, Dict, List, Optional, Type, TypeVar, Protocol, runtime_checkable from dataclasses import dataclass, field from enum import Enum import time T = TypeVar("T") @runtime_checkable class Registerable(Protocol[T]): """可注册对象的协议""" name: str description: str = "" @dataclass class RegistryEntry: """注册条目""" name: str instance: Any metadata: Dict[str, Any] = field(default_factory=dict) registered_at: float = field(default_factory=time.time) class BaseRegistry: """基础注册器""" _entries: Dict[str, RegistryEntry] = {} @classmethod def register(cls, name: str, metadata: Optional[Dict[str, Any]] = None): """ 注册装饰器 Usage: @NodeRegistry.register("my_node", metadata={"category": "processing"}) def my_node(state: AgentState) -> AgentState: ... """ def decorator(obj: T) -> T: entry = RegistryEntry( name=name, instance=obj, metadata=metadata or {}, ) cls._entries[name] = entry return obj return decorator @classmethod def get(cls, name: str) -> Optional[Any]: """获取注册的对象""" entry = cls._entries.get(name) return entry.instance if entry else None @classmethod def get_entry(cls, name: str) -> Optional[RegistryEntry]: """获取注册条目(包含元数据)""" return cls._entries.get(name) @classmethod def list_names(cls) -> List[str]: """列出所有注册名称""" return list(cls._entries.keys()) @classmethod def list_entries(cls) -> List[RegistryEntry]: """列出所有注册条目""" return list(cls._entries.values()) @classmethod def unregister(cls, name: str) -> bool: """注销注册""" if name in cls._entries: del cls._entries[name] return True return False @classmethod def clear(cls): """清空注册表""" cls._entries.clear() class NodeRegistry(BaseRegistry): """节点注册器 - 用于注册 Agent 工作流节点""" _entries: Dict[str, RegistryEntry] = {} class ToolRegistry(BaseRegistry): """工具注册器 - 用于注册工具""" _entries: Dict[str, RegistryEntry] = {} class WorkflowRegistry(BaseRegistry): """工作流注册器 - 用于注册工作流类型""" _entries: Dict[str, RegistryEntry] = {} class ProviderRegistry(BaseRegistry): """LLM Provider 注册器""" _entries: Dict[str, RegistryEntry] = {} @dataclass class ToolMetadata: """工具元数据""" name: str description: str version: str = "1.0.0" timeout: int = 30 retry: int = 0 parameters_schema: Optional[Dict[str, Any]] = None requires_auth: bool = False tags: List[str] = field(default_factory=list) @dataclass class WorkflowMetadata: """工作流元数据""" name: str description: str version: str = "1.0.0" agent_class: Optional[Type] = None default_model: Optional[str] = None supported_features: List[str] = field(default_factory=list) @dataclass class ProviderMetadata: """LLM Provider 元数据""" name: str description: str provider_type: str supported_models: List[str] = field(default_factory=list) config_schema: Optional[Dict[str, Any]] = None def register_tool( name: str, description: str = "", version: str = "1.0.0", timeout: int = 30, retry: int = 0, tags: Optional[List[str]] = None, ): """ 工具注册装饰器 Usage: @register_tool("calculator", "数学计算", timeout=10) class CalculatorTool(BaseTool): ... """ metadata = ToolMetadata( name=name, description=description, version=version, timeout=timeout, retry=retry, tags=tags or [], ) return ToolRegistry.register(name, {"tool_metadata": metadata}) def register_workflow( name: str, description: str = "", version: str = "1.0.0", default_model: Optional[str] = None, supported_features: Optional[List[str]] = None, ): """ 工作流注册装饰器 Usage: @register_workflow("data_query", "数据查询工作流") class DataQueryAgent(BaseAgent): ... """ metadata = WorkflowMetadata( name=name, description=description, version=version, default_model=default_model, supported_features=supported_features or [], ) return WorkflowRegistry.register(name, {"workflow_metadata": metadata}) def register_provider( name: str, provider_type: str, description: str = "", supported_models: Optional[List[str]] = None, ): """ LLM Provider 注册装饰器 Usage: @register_provider("openai", "openai", supported_models=["gpt-4", "gpt-3.5"]) class OpenAIProvider: ... """ metadata = ProviderMetadata( name=name, description=description, provider_type=provider_type, supported_models=supported_models or [], ) return ProviderRegistry.register(name, {"provider_metadata": metadata})