Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
test_provider_manager.py184 linesDownload Raw Back to core
1from core.entities.provider_entities import ModelSettings2from core.model_runtime.entities.model_entities import ModelType3from core.model_runtime.model_providers import model_provider_factory4from core.provider_manager import ProviderManager5from models.provider import LoadBalancingModelConfig, ProviderModelSetting6 7 8def test__to_model_settings(mocker):9    # Get all provider entities10    provider_entities = model_provider_factory.get_providers()11 12    provider_entity = None13    for provider in provider_entities:14        if provider.provider == "openai":15            provider_entity = provider16 17    # Mocking the inputs18    provider_model_settings = [19        ProviderModelSetting(20            id="id",21            tenant_id="tenant_id",22            provider_name="openai",23            model_name="gpt-4",24            model_type="text-generation",25            enabled=True,26            load_balancing_enabled=True,27        )28    ]29    load_balancing_model_configs = [30        LoadBalancingModelConfig(31            id="id1",32            tenant_id="tenant_id",33            provider_name="openai",34            model_name="gpt-4",35            model_type="text-generation",36            name="__inherit__",37            encrypted_config=None,38            enabled=True,39        ),40        LoadBalancingModelConfig(41            id="id2",42            tenant_id="tenant_id",43            provider_name="openai",44            model_name="gpt-4",45            model_type="text-generation",46            name="first",47            encrypted_config='{"openai_api_key": "fake_key"}',48            enabled=True,49        ),50    ]51 52    mocker.patch(53        "core.helper.model_provider_cache.ProviderCredentialsCache.get", return_value={"openai_api_key": "fake_key"}54    )55 56    provider_manager = ProviderManager()57 58    # Running the method59    result = provider_manager._to_model_settings(provider_entity, provider_model_settings, load_balancing_model_configs)60 61    # Asserting that the result is as expected62    assert len(result) == 163    assert isinstance(result[0], ModelSettings)64    assert result[0].model == "gpt-4"65    assert result[0].model_type == ModelType.LLM66    assert result[0].enabled is True67    assert len(result[0].load_balancing_configs) == 268    assert result[0].load_balancing_configs[0].name == "__inherit__"69    assert result[0].load_balancing_configs[1].name == "first"70 71 72def test__to_model_settings_only_one_lb(mocker):73    # Get all provider entities74    provider_entities = model_provider_factory.get_providers()75 76    provider_entity = None77    for provider in provider_entities:78        if provider.provider == "openai":79            provider_entity = provider80 81    # Mocking the inputs82    provider_model_settings = [83        ProviderModelSetting(84            id="id",85            tenant_id="tenant_id",86            provider_name="openai",87            model_name="gpt-4",88            model_type="text-generation",89            enabled=True,90            load_balancing_enabled=True,91        )92    ]93    load_balancing_model_configs = [94        LoadBalancingModelConfig(95            id="id1",96            tenant_id="tenant_id",97            provider_name="openai",98            model_name="gpt-4",99            model_type="text-generation",100            name="__inherit__",101            encrypted_config=None,102            enabled=True,103        )104    ]105 106    mocker.patch(107        "core.helper.model_provider_cache.ProviderCredentialsCache.get", return_value={"openai_api_key": "fake_key"}108    )109 110    provider_manager = ProviderManager()111 112    # Running the method113    result = provider_manager._to_model_settings(provider_entity, provider_model_settings, load_balancing_model_configs)114 115    # Asserting that the result is as expected116    assert len(result) == 1117    assert isinstance(result[0], ModelSettings)118    assert result[0].model == "gpt-4"119    assert result[0].model_type == ModelType.LLM120    assert result[0].enabled is True121    assert len(result[0].load_balancing_configs) == 0122 123 124def test__to_model_settings_lb_disabled(mocker):125    # Get all provider entities126    provider_entities = model_provider_factory.get_providers()127 128    provider_entity = None129    for provider in provider_entities:130        if provider.provider == "openai":131            provider_entity = provider132 133    # Mocking the inputs134    provider_model_settings = [135        ProviderModelSetting(136            id="id",137            tenant_id="tenant_id",138            provider_name="openai",139            model_name="gpt-4",140            model_type="text-generation",141            enabled=True,142            load_balancing_enabled=False,143        )144    ]145    load_balancing_model_configs = [146        LoadBalancingModelConfig(147            id="id1",148            tenant_id="tenant_id",149            provider_name="openai",150            model_name="gpt-4",151            model_type="text-generation",152            name="__inherit__",153            encrypted_config=None,154            enabled=True,155        ),156        LoadBalancingModelConfig(157            id="id2",158            tenant_id="tenant_id",159            provider_name="openai",160            model_name="gpt-4",161            model_type="text-generation",162            name="first",163            encrypted_config='{"openai_api_key": "fake_key"}',164            enabled=True,165        ),166    ]167 168    mocker.patch(169        "core.helper.model_provider_cache.ProviderCredentialsCache.get", return_value={"openai_api_key": "fake_key"}170    )171 172    provider_manager = ProviderManager()173 174    # Running the method175    result = provider_manager._to_model_settings(provider_entity, provider_model_settings, load_balancing_model_configs)176 177    # Asserting that the result is as expected178    assert len(result) == 1179    assert isinstance(result[0], ModelSettings)180    assert result[0].model == "gpt-4"181    assert result[0].model_type == ModelType.LLM182    assert result[0].enabled is True183    assert len(result[0].load_balancing_configs) == 0184