50 lines
1.6 KiB
Python
50 lines
1.6 KiB
Python
"""Protocol adapter contract for external model APIs."""
|
|
|
|
from typing import Any
|
|
|
|
from agenteval.models import ModelCapability, ModelProtocol
|
|
|
|
|
|
class ProtocolAdapterError(ValueError):
|
|
pass
|
|
|
|
|
|
class ModelProtocolAdapter:
|
|
protocol: ModelProtocol
|
|
supported_capabilities: frozenset[ModelCapability] = frozenset()
|
|
|
|
def headers(self, api_key: str | None) -> dict[str, str]:
|
|
return {"Content-Type": "application/json"}
|
|
|
|
def chat_payload(
|
|
self,
|
|
model_name: str | None,
|
|
messages: list[dict[str, str]],
|
|
temperature: float,
|
|
) -> dict[str, Any]:
|
|
self._unsupported(ModelCapability.CHAT)
|
|
|
|
def parse_chat(self, data: dict[str, Any]) -> str:
|
|
self._unsupported(ModelCapability.CHAT)
|
|
|
|
def embedding_payload(self, model_name: str | None, inputs: str | list[str]) -> dict[str, Any]:
|
|
self._unsupported(ModelCapability.EMBEDDING)
|
|
|
|
def parse_embeddings(self, data: dict[str, Any]) -> list[list[float]]:
|
|
self._unsupported(ModelCapability.EMBEDDING)
|
|
|
|
def moderation_payload(self, model_name: str | None, text: str) -> dict[str, Any]:
|
|
self._unsupported(ModelCapability.MODERATION)
|
|
|
|
def parse_moderation(self, data: dict[str, Any]) -> dict[str, Any]:
|
|
self._unsupported(ModelCapability.MODERATION)
|
|
|
|
def _unsupported(self, capability: ModelCapability) -> None:
|
|
raise ProtocolAdapterError(f"{self.protocol.value} 协议不支持 {capability.value} 能力")
|
|
|
|
@staticmethod
|
|
def require_model(model_name: str | None) -> str:
|
|
if not model_name:
|
|
raise ProtocolAdapterError("模型名称不能为空")
|
|
return model_name
|