252 lines
9.3 KiB
Python
252 lines
9.3 KiB
Python
"""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}",
|
||
)
|