Files
more_dots/core/registry.py
T

223 lines
5.5 KiB
Python
Raw Normal View History

2026-03-11 23:40:39 +08:00
"""
核心注册机制模块
提供统一的注册器模式,支持动态扩展:
- 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})