223 lines
5.5 KiB
Python
223 lines
5.5 KiB
Python
"""
|
|
核心注册机制模块
|
|
|
|
提供统一的注册器模式,支持动态扩展:
|
|
- 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})
|