diff --git a/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py b/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py index bff91d2aea3..b5489c88a03 100644 --- a/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py +++ b/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py @@ -1,12 +1,16 @@ from __future__ import annotations +import json import logging +from collections.abc import Iterator from contextlib import contextmanager from types import SimpleNamespace from typing import Any -from unittest.mock import Mock, patch +from unittest.mock import Mock, PropertyMock, patch import pytest +from sqlalchemy import Engine, event, select +from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from core.entities.model_entities import ModelStatus @@ -26,6 +30,7 @@ from core.entities.provider_entities import ( SystemConfigurationStatus, ) from core.helper.model_provider_cache import ProviderCredentialsCacheType +from extensions.ext_database import db from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelType from graphon.model_runtime.entities.provider_entities import ( @@ -38,7 +43,16 @@ from graphon.model_runtime.entities.provider_entities import ( ProviderEntity, ) from models.enums import CredentialSourceType -from models.provider import ProviderType +from models.provider import ( + LoadBalancingModelConfig, + Provider, + ProviderCredential, + ProviderModel, + ProviderModelCredential, + ProviderModelSetting, + ProviderType, + TenantPreferredModelProvider, +) from models.provider_ids import ModelProviderID _UNSET = object() @@ -2158,3 +2172,768 @@ def test_get_custom_provider_models_skips_custom_models_on_schema_error_or_none( assert "get custom model schema failed, boom" in caplog.messages assert any(model.model == "ok-custom" for model in models) assert all(model.model != "none-custom" for model in models) + + +@pytest.fixture +def sqlite_provider_session( + sqlite_session: Session, + sqlite_engine: Engine, +) -> Iterator[Session]: + """Bind provider-owned sessions to the same isolated SQLite database as the test.""" + with patch.object(type(db), "engine", new_callable=PropertyMock, return_value=sqlite_engine): + yield sqlite_session + + +def _provider_credential( + session: Session, + *, + name: str = "API KEY 1", + tenant_id: str = "tenant-1", + provider_name: str = "openai", + encrypted_config: str = "{}", +) -> ProviderCredential: + record = ProviderCredential( + tenant_id=tenant_id, + provider_name=provider_name, + credential_name=name, + encrypted_config=encrypted_config, + ) + session.add(record) + session.commit() + return record + + +def _provider_record( + session: Session, + *, + credential_id: str | None = None, + tenant_id: str = "tenant-1", + provider_name: str = "openai", +) -> Provider: + record = Provider( + tenant_id=tenant_id, + provider_name=provider_name, + provider_type=ProviderType.CUSTOM, + credential_id=credential_id, + is_valid=True, + ) + session.add(record) + session.commit() + return record + + +def _model_credential( + session: Session, + *, + name: str = "API KEY 1", + tenant_id: str = "tenant-1", + provider_name: str = "openai", + model: str = "gpt-4o", + encrypted_config: str = "{}", +) -> ProviderModelCredential: + record = ProviderModelCredential( + tenant_id=tenant_id, + provider_name=provider_name, + model_name=model, + model_type=ModelType.LLM, + credential_name=name, + encrypted_config=encrypted_config, + ) + session.add(record) + session.commit() + return record + + +def _provider_model_record( + session: Session, + *, + credential_id: str | None = None, + tenant_id: str = "tenant-1", + provider_name: str = "openai", + model: str = "gpt-4o", +) -> ProviderModel: + record = ProviderModel( + tenant_id=tenant_id, + provider_name=provider_name, + model_name=model, + model_type=ModelType.LLM, + credential_id=credential_id, + is_valid=True, + ) + session.add(record) + session.commit() + return record + + +def _load_balancing_config( + session: Session, + *, + credential_id: str, + source: CredentialSourceType, + name: str = "Old", +) -> LoadBalancingModelConfig: + record = LoadBalancingModelConfig( + tenant_id="tenant-1", + provider_name="openai", + model_name="gpt-4o", + model_type=ModelType.LLM, + name=name, + encrypted_config="{}", + credential_id=credential_id, + credential_source_type=source, + ) + session.add(record) + session.commit() + return record + + +@contextmanager +def _raise_on_sql(engine: Engine, table_name: str, operation: str) -> Iterator[None]: + """Fail one table operation while production still owns a real transaction.""" + + def fail_target(_conn, _cursor, statement, _parameters, _context, _executemany): + if statement.lstrip().upper().startswith(operation) and table_name in statement: + raise RuntimeError(f"forced {operation} failure for {table_name}") + + event.listen(engine, "before_cursor_execute", fail_target) + try: + yield + finally: + event.remove(engine, "before_cursor_execute", fail_target) + + +@contextmanager +def _mock_cache_boundaries() -> Iterator[tuple[Mock, Mock]]: + with ( + patch("core.entities.provider_configuration.ProviderCredentialsCache") as credentials_cache, + patch.object(ProviderConfiguration, "_invalidate_provider_configuration_cache") as configuration_cache, + ): + yield credentials_cache, configuration_cache + + +def test_generate_credential_names_from_real_rows_and_tenant_isolation( + sqlite_provider_session: Session, +) -> None: + configuration = _build_provider_configuration() + _provider_credential(sqlite_provider_session, name="API KEY 9") + _provider_credential(sqlite_provider_session, name="legacy") + _provider_credential(sqlite_provider_session, name="API KEY 50", tenant_id="other-tenant") + _model_credential(sqlite_provider_session, name="API KEY 4") + assert configuration._generate_provider_credential_name(sqlite_provider_session) == "API KEY 10" + assert ( + configuration._generate_custom_model_credential_name("gpt-4o", ModelType.LLM, sqlite_provider_session) + == "API KEY 5" + ) + + +def test_validate_provider_credentials_reuses_hidden_secret(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + configuration.provider.provider_credential_schema = _build_secret_provider_schema() + credential = _provider_credential(sqlite_provider_session, encrypted_config='{"openai_api_key":"enc-old"}') + factory = Mock() + factory.provider_credentials_validate.return_value = {"openai_api_key": "raw"} + with ( + patch( + "core.entities.provider_configuration.create_plugin_model_assembly", + return_value=SimpleNamespace(model_runtime=Mock(), model_provider_factory=factory), + ), + patch("core.entities.provider_configuration.encrypter.decrypt_token", return_value="raw"), + patch("core.entities.provider_configuration.encrypter.encrypt_token", return_value="enc-new"), + ): + result = configuration.validate_provider_credentials( + {"openai_api_key": HIDDEN_VALUE}, credential_id=credential.id + ) + assert result == {"openai_api_key": "enc-new"} + + +def test_preferred_provider_state_updates_and_is_tenant_scoped(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + configuration.preferred_provider_type = ProviderType.CUSTOM + other = TenantPreferredModelProvider( + tenant_id="other-tenant", provider_name="openai", preferred_provider_type=ProviderType.CUSTOM + ) + current = TenantPreferredModelProvider( + tenant_id="tenant-1", provider_name="openai", preferred_provider_type=ProviderType.CUSTOM + ) + sqlite_provider_session.add_all([other, current]) + sqlite_provider_session.commit() + assert configuration.switch_preferred_provider_type(ProviderType.SYSTEM, session=sqlite_provider_session) + sqlite_provider_session.refresh(current) + sqlite_provider_session.refresh(other) + assert current.preferred_provider_type == ProviderType.SYSTEM + assert other.preferred_provider_type == ProviderType.CUSTOM + + +def test_provider_record_duplicate_and_setting_helpers_use_real_session(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + provider = _provider_record(sqlite_provider_session) + _provider_record(sqlite_provider_session, tenant_id="other-tenant") + credential = _provider_credential(sqlite_provider_session, name="Main") + _provider_credential(sqlite_provider_session, name="Main", tenant_id="other-tenant") + setting = ProviderModelSetting( + tenant_id="tenant-1", + provider_name="openai", + model_name="gpt-4o", + model_type=ModelType.LLM, + ) + sqlite_provider_session.add(setting) + sqlite_provider_session.commit() + assert configuration._get_provider_record(sqlite_provider_session).id == provider.id + assert configuration._check_provider_credential_name_exists("Main", sqlite_provider_session) + assert not configuration._check_provider_credential_name_exists( + "Main", sqlite_provider_session, exclude_id=credential.id + ) + assert configuration._get_provider_model_setting(ModelType.LLM, "gpt-4o", sqlite_provider_session).id == setting.id + + +def test_create_provider_credential_persists_provider_and_rejects_duplicate( + sqlite_provider_session: Session, +) -> None: + configuration = _build_provider_configuration() + with ( + patch.object(ProviderConfiguration, "validate_provider_credentials", return_value={"api_key": "enc"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + ): + configuration.create_provider_credential({"api_key": "raw"}, "Main") + credential = sqlite_provider_session.scalar( + select(ProviderCredential).where(ProviderCredential.credential_name == "Main") + ) + provider = sqlite_provider_session.scalar(select(Provider).where(Provider.tenant_id == "tenant-1")) + assert credential is not None + assert provider is not None + assert provider.credential_id == credential.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=provider.id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_called_once_with( + preferred_model_providers=True, + provider_credentials=True, + ) + with pytest.raises(ValueError, match="already exists"): + configuration.create_provider_credential({"api_key": "raw"}, "Main") + + +def test_update_provider_credential_propagates_to_load_balancing(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + credential = _provider_credential(sqlite_provider_session, name="Old") + provider = _provider_record(sqlite_provider_session, credential_id=credential.id) + lb_config = _load_balancing_config( + sqlite_provider_session, credential_id=credential.id, source=CredentialSourceType.PROVIDER + ) + with ( + patch.object(ProviderConfiguration, "validate_provider_credentials", return_value={"api_key": "enc-new"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + ): + configuration.update_provider_credential({"api_key": "raw"}, credential.id, "New") + sqlite_provider_session.expire_all() + persisted_credential = sqlite_provider_session.get(ProviderCredential, credential.id) + persisted_lb = sqlite_provider_session.get(LoadBalancingModelConfig, lb_config.id) + assert persisted_credential is not None + assert persisted_credential.credential_name == "New" + assert persisted_lb is not None + assert persisted_lb.name == "New" + assert json.loads(persisted_lb.encrypted_config) == {"api_key": "enc-new"} + assert {cache_call.kwargs["identity_id"] for cache_call in credentials_cache.call_args_list} == { + provider.id, + lb_config.id, + } + assert {cache_call.kwargs["cache_type"] for cache_call in credentials_cache.call_args_list} == { + ProviderCredentialsCacheType.PROVIDER, + ProviderCredentialsCacheType.LOAD_BALANCING_MODEL, + } + assert credentials_cache.return_value.delete.call_count == 2 + configuration_cache.assert_called_once_with( + provider_credentials=True, + provider_load_balancing_configs=True, + ) + + +def test_switch_active_provider_credential_updates_persisted_state_and_cache( + sqlite_provider_session: Session, +) -> None: + configuration = _build_provider_configuration() + configuration.preferred_provider_type = ProviderType.CUSTOM + first = _provider_credential(sqlite_provider_session, name="First") + second = _provider_credential(sqlite_provider_session, name="Second") + provider = _provider_record(sqlite_provider_session, credential_id=first.id) + provider_id = provider.id + + with _mock_cache_boundaries() as (credentials_cache, configuration_cache): + configuration.switch_active_provider_credential(second.id) + + sqlite_provider_session.expire_all() + persisted_provider = sqlite_provider_session.get(Provider, provider_id) + assert persisted_provider is not None + assert persisted_provider.credential_id == second.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=provider_id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_not_called() + + +def test_deleting_active_provider_credential_switches_preference_to_system( + sqlite_provider_session: Session, +) -> None: + configuration = _build_provider_configuration() + configuration.preferred_provider_type = ProviderType.CUSTOM + first = _provider_credential(sqlite_provider_session, name="First") + active = _provider_credential(sqlite_provider_session, name="Active") + provider = _provider_record(sqlite_provider_session, credential_id=active.id) + preferred_provider = TenantPreferredModelProvider( + tenant_id="tenant-1", + provider_name="openai", + preferred_provider_type=ProviderType.CUSTOM, + ) + sqlite_provider_session.add(preferred_provider) + sqlite_provider_session.commit() + first_id = first.id + active_id = active.id + provider_id = provider.id + preferred_provider_id = preferred_provider.id + + with _mock_cache_boundaries() as (credentials_cache, configuration_cache): + configuration.delete_provider_credential(active_id) + + sqlite_provider_session.expire_all() + assert sqlite_provider_session.get(ProviderCredential, active_id) is None + assert sqlite_provider_session.get(ProviderCredential, first_id) is not None + persisted_provider = sqlite_provider_session.get(Provider, provider_id) + persisted_preference = sqlite_provider_session.get(TenantPreferredModelProvider, preferred_provider_id) + assert persisted_provider is not None + assert persisted_provider.credential_id is None + assert persisted_preference is not None + assert persisted_preference.preferred_provider_type == ProviderType.SYSTEM + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=provider_id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_called_once_with( + preferred_model_providers=True, + provider_credentials=True, + provider_load_balancing_configs=False, + ) + + +def test_specific_provider_credential_decrypts_and_obfuscates(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + configuration.provider.provider_credential_schema = _build_secret_provider_schema() + credential = _provider_credential(sqlite_provider_session, encrypted_config='{"openai_api_key":"enc"}') + with ( + patch("core.entities.provider_configuration.encrypter.decrypt_token", return_value="raw"), + patch("core.entities.provider_configuration.encrypter.obfuscated_token", return_value="masked"), + ): + result = configuration._get_specific_provider_credential(credential.id) + assert result == {"openai_api_key": "masked"} + with pytest.raises(ValueError, match="not found"): + configuration._get_specific_provider_credential("missing") + + +def test_validate_custom_model_credentials_reuses_hidden_secret(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + configuration.provider.model_credential_schema = _build_secret_model_schema() + credential = _model_credential(sqlite_provider_session, encrypted_config='{"openai_api_key":"enc-old"}') + factory = Mock() + factory.model_credentials_validate.return_value = {"openai_api_key": "raw"} + with ( + patch( + "core.entities.provider_configuration.create_plugin_model_assembly", + return_value=SimpleNamespace(model_runtime=Mock(), model_provider_factory=factory), + ), + patch("core.entities.provider_configuration.encrypter.decrypt_token", return_value="raw"), + patch("core.entities.provider_configuration.encrypter.encrypt_token", return_value="enc-new"), + ): + result = configuration.validate_custom_model_credentials( + ModelType.LLM, + "gpt-4o", + {"openai_api_key": HIDDEN_VALUE}, + credential_id=credential.id, + ) + assert result == {"openai_api_key": "enc-new"} + + +def test_specific_custom_model_credential_preserves_secret_when_decryption_fails( + sqlite_provider_session: Session, + caplog: pytest.LogCaptureFixture, +) -> None: + configuration = _build_provider_configuration() + configuration.provider.model_credential_schema = _build_secret_model_schema() + credential = _model_credential( + sqlite_provider_session, + name="Main", + encrypted_config='{"openai_api_key":"enc-secret"}', + ) + + with ( + caplog.at_level(logging.ERROR, logger="core.entities.provider_configuration"), + patch("core.entities.provider_configuration.encrypter.decrypt_token", side_effect=RuntimeError("boom")), + patch.object( + ProviderConfiguration, + "obfuscated_credentials", + side_effect=lambda credentials, credential_form_schemas: credentials, + ), + ): + result = configuration._get_specific_custom_model_credential(ModelType.LLM, "gpt-4o", credential.id) + + assert result == { + "current_credential_id": credential.id, + "current_credential_name": "Main", + "credentials": {"openai_api_key": "enc-secret"}, + } + assert caplog.messages.count("Failed to decrypt model credential secret variable openai_api_key") == 1 + + +def test_create_update_and_delete_custom_model_credential(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + with ( + patch.object(ProviderConfiguration, "validate_custom_model_credentials", return_value={"api_key": "enc"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + ): + configuration.create_custom_model_credential(ModelType.LLM, "gpt-4o", {"api_key": "raw"}, "Main") + credential = sqlite_provider_session.scalar(select(ProviderModelCredential)) + model = sqlite_provider_session.scalar(select(ProviderModel)) + assert credential is not None + assert model is not None + assert model.credential_id == credential.id + credential_id = credential.id + model_id = model.id + lb_config = _load_balancing_config( + sqlite_provider_session, + credential_id=credential_id, + source=CredentialSourceType.CUSTOM_MODEL, + ) + lb_config_id = lb_config.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=model_id, + cache_type=ProviderCredentialsCacheType.MODEL, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_called_once_with( + provider_models=True, + provider_model_credentials=True, + ) + + with ( + patch.object(ProviderConfiguration, "validate_custom_model_credentials", return_value={"api_key": "enc-2"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + ): + configuration.update_custom_model_credential( + ModelType.LLM, "gpt-4o", {"api_key": "raw"}, "Renamed", credential_id + ) + sqlite_provider_session.expire_all() + persisted_credential = sqlite_provider_session.get(ProviderModelCredential, credential_id) + persisted_lb = sqlite_provider_session.get(LoadBalancingModelConfig, lb_config_id) + assert persisted_credential is not None + assert persisted_credential.credential_name == "Renamed" + assert json.loads(persisted_credential.encrypted_config) == {"api_key": "enc-2"} + assert persisted_lb is not None + assert persisted_lb.name == "Renamed" + assert json.loads(persisted_lb.encrypted_config) == {"api_key": "enc-2"} + assert {cache_call.kwargs["identity_id"] for cache_call in credentials_cache.call_args_list} == { + model_id, + lb_config_id, + } + assert {cache_call.kwargs["cache_type"] for cache_call in credentials_cache.call_args_list} == { + ProviderCredentialsCacheType.MODEL, + ProviderCredentialsCacheType.LOAD_BALANCING_MODEL, + } + assert credentials_cache.return_value.delete.call_count == 2 + configuration_cache.assert_called_once_with( + provider_models=True, + provider_model_credentials=True, + provider_load_balancing_configs=True, + ) + + with _mock_cache_boundaries() as (credentials_cache, configuration_cache): + configuration.delete_custom_model_credential(ModelType.LLM, "gpt-4o", credential_id) + sqlite_provider_session.expire_all() + assert sqlite_provider_session.get(ProviderModelCredential, credential_id) is None + assert sqlite_provider_session.get(ProviderModel, model_id) is None + assert sqlite_provider_session.get(LoadBalancingModelConfig, lb_config_id) is None + assert {cache_call.kwargs["identity_id"] for cache_call in credentials_cache.call_args_list} == { + model_id, + lb_config_id, + } + assert {cache_call.kwargs["cache_type"] for cache_call in credentials_cache.call_args_list} == { + ProviderCredentialsCacheType.MODEL, + ProviderCredentialsCacheType.LOAD_BALANCING_MODEL, + } + assert credentials_cache.return_value.delete.call_count == 2 + configuration_cache.assert_called_once_with( + provider_models=True, + provider_model_credentials=True, + provider_load_balancing_configs=True, + ) + + +def test_add_and_switch_custom_model_credential(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + first = _model_credential(sqlite_provider_session, name="First") + second = _model_credential(sqlite_provider_session, name="Second") + with _mock_cache_boundaries() as (credentials_cache, configuration_cache): + configuration.add_model_credential_to_model(ModelType.LLM, "gpt-4o", first.id) + configuration.switch_custom_model_credential(ModelType.LLM, "gpt-4o", second.id) + model = sqlite_provider_session.scalar(select(ProviderModel)) + assert model is not None + assert model.credential_id == second.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=model.id, + cache_type=ProviderCredentialsCacheType.MODEL, + ) + credentials_cache.return_value.delete.assert_called_once_with() + assert configuration_cache.call_count == 2 + for cache_call in configuration_cache.call_args_list: + assert cache_call.kwargs == {"provider_models": True} + with pytest.raises(ValueError, match="Can't add same credential"): + configuration.add_model_credential_to_model(ModelType.LLM, "gpt-4o", second.id) + + +def test_model_settings_and_load_balancing_persist(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + with patch.object(configuration, "_invalidate_provider_configuration_cache") as configuration_cache: + configuration.disable_model(ModelType.LLM, "gpt-4o") + configuration_cache.assert_called_once_with(provider_model_settings=True) + persisted_setting = sqlite_provider_session.scalar(select(ProviderModelSetting)) + assert persisted_setting is not None + assert persisted_setting.enabled is False + with patch.object(configuration, "_invalidate_provider_configuration_cache") as configuration_cache: + configuration.enable_model(ModelType.LLM, "gpt-4o") + configuration_cache.assert_called_once_with(provider_model_settings=True) + sqlite_provider_session.expire_all() + refreshed_setting = sqlite_provider_session.get(ProviderModelSetting, persisted_setting.id) + assert refreshed_setting is not None + assert refreshed_setting.enabled is True + + first = _provider_credential(sqlite_provider_session, name="First") + second = _provider_credential(sqlite_provider_session, name="Second") + _load_balancing_config(sqlite_provider_session, credential_id=first.id, source=CredentialSourceType.PROVIDER) + _load_balancing_config(sqlite_provider_session, credential_id=second.id, source=CredentialSourceType.PROVIDER) + with patch.object(configuration, "_invalidate_provider_configuration_cache") as configuration_cache: + configuration.enable_model_load_balancing(ModelType.LLM, "gpt-4o") + configuration_cache.assert_called_once_with(provider_model_settings=True) + sqlite_provider_session.expire_all() + refreshed_setting = sqlite_provider_session.get(ProviderModelSetting, persisted_setting.id) + assert refreshed_setting is not None + assert refreshed_setting.load_balancing_enabled is True + with patch.object(configuration, "_invalidate_provider_configuration_cache") as configuration_cache: + configuration.disable_model_load_balancing(ModelType.LLM, "gpt-4o") + configuration_cache.assert_called_once_with(provider_model_settings=True) + sqlite_provider_session.expire_all() + refreshed_setting = sqlite_provider_session.get(ProviderModelSetting, persisted_setting.id) + assert refreshed_setting is not None + assert refreshed_setting.load_balancing_enabled is False + + +def test_provider_create_rolls_back_on_insert_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + with ( + patch.object(ProviderConfiguration, "validate_provider_credentials", return_value={"api_key": "enc"}), + _mock_cache_boundaries(), + _raise_on_sql(sqlite_engine, "provider_credentials", "INSERT"), + pytest.raises(RuntimeError, match="forced INSERT"), + ): + configuration.create_provider_credential({"api_key": "raw"}, "Main") + assert sqlite_provider_session.scalar(select(ProviderCredential)) is None + assert sqlite_provider_session.scalar(select(Provider)) is None + + +def test_custom_model_create_rolls_back_on_insert_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + with ( + patch.object(ProviderConfiguration, "validate_custom_model_credentials", return_value={"api_key": "enc"}), + _mock_cache_boundaries(), + _raise_on_sql(sqlite_engine, "provider_model_credentials", "INSERT"), + pytest.raises(RuntimeError, match="forced INSERT"), + ): + configuration.create_custom_model_credential(ModelType.LLM, "gpt-4o", {"api_key": "raw"}, "Main") + assert sqlite_provider_session.scalar(select(ProviderModelCredential)) is None + assert sqlite_provider_session.scalar(select(ProviderModel)) is None + + +def test_provider_update_rolls_back_on_update_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + credential = _provider_credential( + sqlite_provider_session, + name="Old", + encrypted_config='{"api_key":"enc-old"}', + ) + _provider_record(sqlite_provider_session, credential_id=credential.id) + + with ( + patch.object(ProviderConfiguration, "validate_provider_credentials", return_value={"api_key": "enc-new"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_credentials", "UPDATE"), + pytest.raises(RuntimeError, match="forced UPDATE"), + ): + configuration.update_provider_credential({"api_key": "raw"}, credential.id, "New") + + sqlite_provider_session.expire_all() + persisted = sqlite_provider_session.get(ProviderCredential, credential.id) + assert persisted is not None + assert persisted.credential_name == "Old" + assert json.loads(persisted.encrypted_config) == {"api_key": "enc-old"} + credentials_cache.assert_not_called() + configuration_cache.assert_not_called() + + +def test_provider_delete_rolls_back_on_delete_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + credential = _provider_credential(sqlite_provider_session) + provider = _provider_record(sqlite_provider_session, credential_id=credential.id) + + with ( + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_credentials", "DELETE"), + pytest.raises(RuntimeError, match="forced DELETE"), + ): + configuration.delete_provider_credential(credential.id) + + sqlite_provider_session.expire_all() + assert sqlite_provider_session.get(ProviderCredential, credential.id) is not None + persisted_provider = sqlite_provider_session.get(Provider, provider.id) + assert persisted_provider is not None + assert persisted_provider.credential_id == credential.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=provider.id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_not_called() + + +def test_provider_switch_rolls_back_on_update_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + first = _provider_credential(sqlite_provider_session, name="First") + second = _provider_credential(sqlite_provider_session, name="Second") + provider = _provider_record(sqlite_provider_session, credential_id=first.id) + + with ( + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "providers", "UPDATE"), + pytest.raises(RuntimeError, match="forced UPDATE"), + ): + configuration.switch_active_provider_credential(second.id) + + sqlite_provider_session.expire_all() + persisted_provider = sqlite_provider_session.get(Provider, provider.id) + assert persisted_provider is not None + assert persisted_provider.credential_id == first.id + credentials_cache.assert_not_called() + configuration_cache.assert_not_called() + + +def test_custom_model_update_rolls_back_on_update_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + credential = _model_credential( + sqlite_provider_session, + name="Old", + encrypted_config='{"api_key":"enc-old"}', + ) + _provider_model_record(sqlite_provider_session, credential_id=credential.id) + + with ( + patch.object( + ProviderConfiguration, + "validate_custom_model_credentials", + return_value={"api_key": "enc-new"}, + ), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_model_credentials", "UPDATE"), + pytest.raises(RuntimeError, match="forced UPDATE"), + ): + configuration.update_custom_model_credential( + ModelType.LLM, + "gpt-4o", + {"api_key": "raw"}, + "New", + credential.id, + ) + + sqlite_provider_session.expire_all() + persisted = sqlite_provider_session.get(ProviderModelCredential, credential.id) + assert persisted is not None + assert persisted.credential_name == "Old" + assert json.loads(persisted.encrypted_config) == {"api_key": "enc-old"} + credentials_cache.assert_not_called() + configuration_cache.assert_not_called() + + +def test_custom_model_delete_rolls_back_on_delete_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + credential = _model_credential(sqlite_provider_session) + model = _provider_model_record(sqlite_provider_session, credential_id=credential.id) + + with ( + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_model_credentials", "DELETE"), + pytest.raises(RuntimeError, match="forced DELETE"), + ): + configuration.delete_custom_model_credential(ModelType.LLM, "gpt-4o", credential.id) + + sqlite_provider_session.expire_all() + assert sqlite_provider_session.get(ProviderModelCredential, credential.id) is not None + persisted_model = sqlite_provider_session.get(ProviderModel, model.id) + assert persisted_model is not None + assert persisted_model.credential_id == credential.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=model.id, + cache_type=ProviderCredentialsCacheType.MODEL, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_not_called() + + +def test_custom_model_switch_rolls_back_on_update_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + first = _model_credential(sqlite_provider_session, name="First") + second = _model_credential(sqlite_provider_session, name="Second") + model = _provider_model_record(sqlite_provider_session, credential_id=first.id) + + with ( + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_models", "UPDATE"), + pytest.raises(RuntimeError, match="forced UPDATE"), + ): + configuration.switch_custom_model_credential(ModelType.LLM, "gpt-4o", second.id) + + sqlite_provider_session.expire_all() + persisted_model = sqlite_provider_session.get(ProviderModel, model.id) + assert persisted_model is not None + assert persisted_model.credential_id == first.id + credentials_cache.assert_not_called() + configuration_cache.assert_not_called()