x
This commit is contained in:
@@ -0,0 +1,222 @@
|
||||
"""
|
||||
核心注册机制模块
|
||||
|
||||
提供统一的注册器模式,支持动态扩展:
|
||||
- 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})
|
||||
Reference in New Issue
Block a user