111 lines
3.8 KiB
Python
111 lines
3.8 KiB
Python
"""Model configuration persistence, encryption, and binding rules."""
|
|
|
|
import pytest
|
|
from agenteval.models import Case, ModelPurpose, Scenario
|
|
from agenteval.services.model_configs import (
|
|
ModelConfigError,
|
|
ModelConfigInUseError,
|
|
ModelConfigService,
|
|
SecretCipher,
|
|
SecretKeyError,
|
|
)
|
|
from agenteval.storage.repository import ScenarioRepository
|
|
from cryptography.fernet import Fernet
|
|
|
|
|
|
def _service(session):
|
|
return ModelConfigService(session, SecretCipher(Fernet.generate_key().decode("ascii")))
|
|
|
|
|
|
def _create_config(service, *, name="评估模型", capability="chat", is_default=False, api_key="secret"):
|
|
return service.create(
|
|
name=name,
|
|
provider="openai_compatible",
|
|
capability=capability,
|
|
endpoint_url="https://models.example.com/v1/chat/completions",
|
|
model_name="test-model" if capability != "moderation" else None,
|
|
api_key=api_key,
|
|
enabled=True,
|
|
is_default=is_default,
|
|
description="test",
|
|
)
|
|
|
|
|
|
def _scenario(config_id: str) -> Scenario:
|
|
return Scenario(
|
|
id="scenario-model-binding",
|
|
name="模型绑定测试",
|
|
cases=[Case(id="case-1", messages=["hello"])],
|
|
model_bindings={ModelPurpose.JUDGE: config_id},
|
|
)
|
|
|
|
|
|
def test_secret_cipher_round_trip_and_invalid_key():
|
|
cipher = SecretCipher(Fernet.generate_key().decode("ascii"))
|
|
encrypted = cipher.encrypt("sk-private")
|
|
assert encrypted and encrypted != "sk-private"
|
|
assert cipher.decrypt(encrypted) == "sk-private"
|
|
|
|
with pytest.raises(SecretKeyError):
|
|
SecretCipher("invalid").encrypt("value")
|
|
with pytest.raises(SecretKeyError):
|
|
SecretCipher("").encrypt("value")
|
|
|
|
|
|
def test_create_update_and_single_default(db_session):
|
|
service = _service(db_session)
|
|
first = _create_config(service, name="first", is_default=True)
|
|
second = _create_config(service, name="second", is_default=True)
|
|
|
|
assert service.require(first.id).is_default is False
|
|
assert service.require(second.id).is_default is True
|
|
encrypted = second.api_key_encrypted
|
|
|
|
updated = service.update(
|
|
second.id,
|
|
name="second",
|
|
provider="openai_compatible",
|
|
capability="chat",
|
|
endpoint_url="https://models.example.com/v1/chat/completions",
|
|
model_name="new-model",
|
|
api_key=None,
|
|
clear_api_key=False,
|
|
enabled=True,
|
|
is_default=True,
|
|
description="updated",
|
|
)
|
|
assert updated.api_key_encrypted == encrypted
|
|
assert service.resolve(second.id).api_key == "secret"
|
|
|
|
|
|
def test_binding_validation_and_reference_protection(db_session):
|
|
service = _service(db_session)
|
|
chat = _create_config(service)
|
|
embedding = _create_config(service, name="向量模型", capability="embedding")
|
|
|
|
with pytest.raises(ModelConfigError, match="不能用于 chat"):
|
|
service.validate_bindings({"judge": embedding.id})
|
|
|
|
created = ScenarioRepository(db_session).create(_scenario(chat.id))
|
|
assert created.model_bindings == {ModelPurpose.JUDGE: chat.id}
|
|
loaded = ScenarioRepository(db_session).get(created.id)
|
|
assert loaded and loaded.model_bindings[ModelPurpose.JUDGE] == chat.id
|
|
|
|
with pytest.raises(ModelConfigInUseError) as exc_info:
|
|
service.delete(chat.id)
|
|
assert exc_info.value.references[0]["scenario_id"] == created.id
|
|
|
|
assert ScenarioRepository(db_session).delete(created.id)
|
|
service.delete(chat.id)
|
|
assert service.repo.get(chat.id) is None
|
|
|
|
|
|
def test_failed_binding_does_not_create_scenario(db_session):
|
|
service = _service(db_session)
|
|
moderation = _create_config(service, name="审核模型", capability="moderation")
|
|
scenario = _scenario(moderation.id)
|
|
|
|
with pytest.raises(ModelConfigError):
|
|
ScenarioRepository(db_session).create(scenario)
|
|
assert ScenarioRepository(db_session).get(scenario.id) is None
|