AgentEvalTool/backend/agenteval/services/model_configs.py

252 lines
9.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Business rules and secret handling for model configurations."""
from dataclasses import dataclass
from datetime import datetime
from urllib.parse import urlparse
from cryptography.fernet import Fernet, InvalidToken
from sqlmodel import Session
from agenteval.config import get_settings
from agenteval.models import ModelCapability, ModelPurpose
from agenteval.storage.db import ModelConfigDB, ScenarioDB
from agenteval.storage.model_config_repository import ModelConfigRepository
PURPOSE_CAPABILITIES: dict[str, ModelCapability] = {
ModelPurpose.GENERATOR.value: ModelCapability.CHAT,
ModelPurpose.JUDGE.value: ModelCapability.CHAT,
ModelPurpose.EMBEDDING.value: ModelCapability.EMBEDDING,
ModelPurpose.MODERATION.value: ModelCapability.MODERATION,
}
class ModelConfigError(ValueError):
pass
class ModelConfigNotFoundError(ModelConfigError):
pass
class ModelConfigInUseError(ModelConfigError):
def __init__(self, references: list[dict[str, str]]):
super().__init__("模型配置正在被场景引用")
self.references = references
class SecretKeyError(ModelConfigError):
pass
@dataclass(frozen=True)
class ModelRuntimeConfig:
id: str
name: str
provider: str
capability: ModelCapability
endpoint_url: str
model_name: str | None
api_key: str | None
updated_at: datetime | None
def snapshot(self) -> dict[str, str | None]:
return {
"id": self.id,
"name": self.name,
"provider": self.provider,
"capability": self.capability.value,
"endpoint_url": self.endpoint_url,
"model_name": self.model_name,
"updated_at": self.updated_at.isoformat() if self.updated_at else None,
}
class SecretCipher:
def __init__(self, key: str | None = None):
self._key = key if key is not None else get_settings().secret_key
def _fernet(self) -> Fernet:
if not self._key:
raise SecretKeyError("未配置 AGENTEVAL_SECRET_KEY无法保存或读取模型 API Key")
try:
return Fernet(self._key.encode("ascii"))
except (ValueError, UnicodeEncodeError) as exc:
raise SecretKeyError("AGENTEVAL_SECRET_KEY 不是有效的 Fernet Key") from exc
def encrypt(self, value: str | None) -> str | None:
if not value:
return None
return self._fernet().encrypt(value.encode("utf-8")).decode("ascii")
def decrypt(self, value: str | None) -> str | None:
if not value:
return None
try:
return self._fernet().decrypt(value.encode("ascii")).decode("utf-8")
except InvalidToken as exc:
raise SecretKeyError("模型 API Key 无法解密,请检查 AGENTEVAL_SECRET_KEY") from exc
class ModelConfigService:
def __init__(self, session: Session, cipher: SecretCipher | None = None):
self.session = session
self.repo = ModelConfigRepository(session)
self.cipher = cipher or SecretCipher()
@staticmethod
def validate_fields(capability: str, endpoint_url: str, model_name: str | None) -> None:
try:
capability_value = ModelCapability(capability)
except ValueError as exc:
raise ModelConfigError(f"不支持的模型能力: {capability}") from exc
parsed = urlparse(endpoint_url)
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
raise ModelConfigError("Endpoint 必须是有效的 HTTP 或 HTTPS URL")
if capability_value in {ModelCapability.CHAT, ModelCapability.EMBEDDING} and not model_name:
raise ModelConfigError("chat 和 embedding 配置必须填写模型名称")
@staticmethod
def validate_common(name: str, provider: str, enabled: bool, is_default: bool) -> None:
if not name.strip():
raise ModelConfigError("模型配置名称不能为空")
if provider != "openai_compatible":
raise ModelConfigError(f"不支持的模型协议: {provider}")
if is_default and not enabled:
raise ModelConfigError("停用的模型配置不能设为默认")
def create(
self,
*,
name: str,
provider: str,
capability: str,
endpoint_url: str,
model_name: str | None,
api_key: str | None,
enabled: bool,
is_default: bool,
description: str,
) -> ModelConfigDB:
name = name.strip()
self.validate_common(name, provider, enabled, is_default)
self.validate_fields(capability, endpoint_url.strip(), model_name)
if self.repo.get_by_name(name):
raise ModelConfigError("模型配置名称已存在")
config = ModelConfigDB(
name=name,
provider=provider,
capability=capability,
endpoint_url=endpoint_url.strip(),
model_name=model_name.strip() if model_name else None,
api_key_encrypted=self.cipher.encrypt(api_key),
enabled=enabled,
is_default=is_default,
description=description.strip(),
)
return self.repo.create(config)
def update(
self,
config_id: str,
*,
name: str,
provider: str,
capability: str,
endpoint_url: str,
model_name: str | None,
api_key: str | None,
clear_api_key: bool,
enabled: bool,
is_default: bool,
description: str,
) -> ModelConfigDB:
config = self.require(config_id)
name = name.strip()
self.validate_common(name, provider, enabled, is_default)
self.validate_fields(capability, endpoint_url.strip(), model_name)
references = self.repo.list_references(config_id)
if references and not enabled:
raise ModelConfigError("模型配置正在被场景引用,不能停用")
for reference in references:
expected = PURPOSE_CAPABILITIES.get(reference.purpose)
if expected and expected.value != capability:
raise ModelConfigError(
f"模型配置正用于 {reference.purpose},能力不能改为 {capability}",
)
same_name = self.repo.get_by_name(name)
if same_name and same_name.id != config_id:
raise ModelConfigError("模型配置名称已存在")
config.name = name
config.provider = provider
config.capability = capability
config.endpoint_url = endpoint_url.strip()
config.model_name = model_name.strip() if model_name else None
config.enabled = enabled
config.is_default = is_default
config.description = description.strip()
if clear_api_key:
config.api_key_encrypted = None
elif api_key:
config.api_key_encrypted = self.cipher.encrypt(api_key)
return self.repo.update(config)
def require(self, config_id: str) -> ModelConfigDB:
config = self.repo.get(config_id)
if not config:
raise ModelConfigNotFoundError("模型配置不存在")
return config
def resolve(
self,
config_id: str,
expected_capability: ModelCapability | None = None,
) -> ModelRuntimeConfig:
config = self.require(config_id)
capability = ModelCapability(config.capability)
if not config.enabled:
raise ModelConfigError(f"模型配置“{config.name}”已禁用")
if expected_capability and capability != expected_capability:
raise ModelConfigError(
f"模型配置“{config.name}”能力为 {capability.value},不能用于 {expected_capability.value}",
)
return ModelRuntimeConfig(
id=config.id or "",
name=config.name,
provider=config.provider,
capability=capability,
endpoint_url=config.endpoint_url,
model_name=config.model_name,
api_key=self.cipher.decrypt(config.api_key_encrypted),
updated_at=config.updated_at,
)
def delete(self, config_id: str) -> None:
config = self.require(config_id)
references = []
for binding in self.repo.list_references(config_id):
scenario = self.session.get(ScenarioDB, binding.scenario_id)
references.append(
{
"scenario_id": binding.scenario_id,
"scenario_name": scenario.name if scenario else binding.scenario_id,
"purpose": binding.purpose,
}
)
if references:
raise ModelConfigInUseError(references)
self.repo.delete(config)
def validate_bindings(self, bindings: dict[str, str]) -> None:
for purpose, config_id in bindings.items():
expected = PURPOSE_CAPABILITIES.get(purpose)
if expected is None:
raise ModelConfigError(f"不支持的模型用途: {purpose}")
config = self.require(config_id)
if not config.enabled:
raise ModelConfigError(f"模型配置“{config.name}”已禁用")
capability = ModelCapability(config.capability)
if capability != expected:
raise ModelConfigError(
f"模型配置“{config.name}”能力为 {capability.value},不能用于 {expected.value}",
)