40 lines
1.1 KiB
Python
40 lines
1.1 KiB
Python
"""Registry of supported external model protocols."""
|
|
|
|
from agenteval.models import ModelCapability, ModelProtocol
|
|
|
|
from .anthropic import AnthropicAdapter
|
|
from .base import ModelProtocolAdapter, ProtocolAdapterError
|
|
from .dashscope import DashScopeAdapter
|
|
from .gemini import GoogleGeminiAdapter
|
|
from .openai import OpenAICompatibleAdapter
|
|
|
|
_ADAPTERS: dict[ModelProtocol, ModelProtocolAdapter] = {
|
|
adapter.protocol: adapter
|
|
for adapter in (
|
|
OpenAICompatibleAdapter(),
|
|
AnthropicAdapter(),
|
|
GoogleGeminiAdapter(),
|
|
DashScopeAdapter(),
|
|
)
|
|
}
|
|
|
|
|
|
def get_protocol_adapter(provider: str | ModelProtocol) -> ModelProtocolAdapter:
|
|
try:
|
|
protocol = ModelProtocol(provider)
|
|
return _ADAPTERS[protocol]
|
|
except (ValueError, KeyError) as exc:
|
|
raise ProtocolAdapterError(f"不支持的模型协议: {provider}") from exc
|
|
|
|
|
|
def get_protocol_capabilities(provider: str | ModelProtocol) -> frozenset[ModelCapability]:
|
|
return get_protocol_adapter(provider).supported_capabilities
|
|
|
|
|
|
__all__ = [
|
|
"ModelProtocolAdapter",
|
|
"ProtocolAdapterError",
|
|
"get_protocol_adapter",
|
|
"get_protocol_capabilities",
|
|
]
|