From e5c3e006c23097fbb5db49e0b13e303ac3e01da4 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 20 Aug 2025 16:33:31 -0700 Subject: [PATCH] fix(proxy_server.py): fix merge config + db router settings logic --- litellm/proxy/proxy_server.py | 28 +- litellm/utils.py | 26 +- tests/test_litellm/proxy/test_proxy_server.py | 573 +++++++++++++----- 3 files changed, 457 insertions(+), 170 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 591747237e7..bccca574b13 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2675,12 +2675,30 @@ class ProxyConfig: db_router_settings = await prisma_client.db.litellm_config.find_first( where={"param_name": "router_settings"} ) - if ( - db_router_settings is not None - and db_router_settings.param_value is not None + + config_router_settings = config_data.get("router_settings", {}) + + combined_router_settings = {} + if config_router_settings is not None and isinstance( + config_router_settings, dict + ) and db_router_settings is not None and isinstance( + db_router_settings.param_value, dict ): - _router_settings = db_router_settings.param_value - llm_router.update_settings(**_router_settings) + from litellm.utils import _update_dictionary + combined_router_settings = _update_dictionary( + config_router_settings, db_router_settings.param_value + ) + elif config_router_settings is not None and isinstance( + config_router_settings, dict + ): + combined_router_settings = config_router_settings + elif db_router_settings is not None and isinstance( + db_router_settings.param_value, dict + ): + combined_router_settings = db_router_settings.param_value + + if combined_router_settings is not None: + llm_router.update_settings(**combined_router_settings) def _add_general_settings_from_db_config( self, config_data: dict, general_settings: dict, proxy_logging_obj: ProxyLogging diff --git a/litellm/utils.py b/litellm/utils.py index 05055b7df8b..8d7360b4fb2 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2255,7 +2255,17 @@ def _update_dictionary(existing_dict: Dict, new_dict: dict) -> dict: for k, v in new_dict.items(): if v is not None: # Convert stringified numbers to appropriate numeric types - existing_dict[k] = _convert_stringified_numbers(v) + if isinstance(v, str): + existing_dict[k] = _convert_stringified_numbers(v) + elif isinstance(v, dict): + existing_nested_dict = existing_dict.get(k) + if isinstance(existing_nested_dict, dict): + existing_nested_dict.update(v) + existing_dict[k] = existing_nested_dict + else: + existing_dict[k] = v + else: + existing_dict[k] = v return existing_dict @@ -2928,19 +2938,19 @@ def _remove_strict_from_schema(schema): def _remove_json_schema_refs(schema, max_depth=10): """ Remove JSON schema reference fields like '$id' and '$schema' that can cause issues with some providers. - + These fields are used for schema validation but can cause problems when the schema references are not accessible to the provider's validation system. - + Args: schema: The schema object to clean (dict, list, or other) max_depth: Maximum recursion depth to prevent infinite loops (default: 10) - + Relevant Issues: Mistral API grammar validation fails when schema contains $id and $schema references """ if max_depth <= 0: return schema - + if isinstance(schema, dict): # Remove JSON schema reference fields schema.pop("$id", None) @@ -6664,7 +6674,8 @@ def validate_and_fix_openai_messages(messages: List): new_messages.append(cleaned_message) return validate_chat_completion_user_messages(messages=new_messages) -def validate_and_fix_openai_tools(tools: Optional[List]) -> Optional[List[dict]]: + +def validate_and_fix_openai_tools(tools: Optional[List]) -> Optional[List[dict]]: """ Ensure tools is List[dict] and not List[BaseModel] """ @@ -6678,6 +6689,7 @@ def validate_and_fix_openai_tools(tools: Optional[List]) -> Optional[List[dict]] new_tools.append(tool) return new_tools + def cleanup_none_field_in_message(message: AllMessageValues): """ Cleans up the message by removing the none field. @@ -7072,6 +7084,7 @@ class ProviderConfigManager: # This mapping ensures that the correct configuration is returned for BEDROCK. elif litellm.LlmProviders.BEDROCK == provider: from litellm.llms.bedrock.common_utils import BedrockModelInfo + return BedrockModelInfo.get_bedrock_provider_config_for_messages_api(model) elif litellm.LlmProviders.VERTEX_AI == provider: if "claude" in model: @@ -7143,6 +7156,7 @@ class ProviderConfigManager: return litellm.GeminiModelInfo() elif LlmProviders.VERTEX_AI == provider: from litellm.llms.vertex_ai.common_utils import VertexAIModelInfo + return VertexAIModelInfo() elif LlmProviders.LITELLM_PROXY == provider: return litellm.LiteLLMProxyChatConfig() diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index c774602e5ca..28410edd533 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -21,8 +21,8 @@ sys.path.insert( ) # Adds the parent directory to the system-path import litellm -from litellm.proxy.proxy_server import app, initialize from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.proxy_server import app, initialize example_embedding_result = { "object": "list", @@ -1119,9 +1119,11 @@ async def test_chat_completion_result_no_nested_none_values(): ) mock_model_response.choices = [mock_choice] - setattr(mock_model_response, "usage", litellm.Usage( - prompt_tokens=10, completion_tokens=5, total_tokens=15 - )) + setattr( + mock_model_response, + "usage", + litellm.Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + ) # Verify the mock has None values before serialization raw_dict = mock_model_response.model_dump() @@ -1198,37 +1200,42 @@ async def test_chat_completion_result_no_nested_none_values(): # Price Data Reload Tests # ============================================================================ + class TestPriceDataReloadAPI: """Test cases for price data reload API endpoints""" - + @pytest.fixture def client_with_auth(self): """Create a test client with authentication""" - from litellm.proxy.proxy_server import cleanup_router_config_variables from litellm.proxy._types import LitellmUserRoles - + from litellm.proxy.proxy_server import cleanup_router_config_variables + cleanup_router_config_variables() filepath = os.path.dirname(os.path.abspath(__file__)) config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml" asyncio.run(initialize(config=config_fp, debug=True)) - + # Mock admin user authentication mock_auth = MagicMock() mock_auth.user_role = LitellmUserRoles.PROXY_ADMIN app.dependency_overrides[user_api_key_auth] = lambda: mock_auth - + return TestClient(app) - + def test_reload_model_cost_map_admin_access(self, client_with_auth): """Test that admin users can access the reload endpoint""" - with patch('litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map') as mock_get_map: - mock_get_map.return_value = {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}} + with patch( + "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map" + ) as mock_get_map: + mock_get_map.return_value = { + "gpt-3.5-turbo": {"input_cost_per_token": 0.001} + } # Mock the database connection - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) - + response = client_with_auth.post("/reload/model_cost_map") - + assert response.status_code == 200 data = response.json() assert data["status"] == "success" @@ -1236,69 +1243,78 @@ class TestPriceDataReloadAPI: assert "timestamp" in data assert "models_count" in data # The new implementation immediately reloads and returns the count - assert "Price data reloaded successfully! 1 models updated." in data["message"] + assert ( + "Price data reloaded successfully! 1 models updated." + in data["message"] + ) assert data["models_count"] == 1 - + def test_reload_model_cost_map_non_admin_access(self, client_with_auth): """Test that non-admin users cannot access the reload endpoint""" # Mock non-admin user mock_auth = MagicMock() mock_auth.user_role = "user" # Non-admin role app.dependency_overrides[user_api_key_auth] = lambda: mock_auth - + response = client_with_auth.post("/reload/model_cost_map") - + assert response.status_code == 403 data = response.json() assert "Access denied" in data["detail"] assert "Admin role required" in data["detail"] - + def test_get_model_cost_map_admin_access(self, client_with_auth): """Test that admin users can access the get model cost map endpoint""" - with patch('litellm.model_cost', {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}}): + with patch( + "litellm.model_cost", {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}} + ): response = client_with_auth.get("/get/litellm_model_cost_map") - + assert response.status_code == 200 data = response.json() assert "gpt-3.5-turbo" in data - + def test_get_model_cost_map_non_admin_access(self, client_with_auth): """Test that non-admin users cannot access the get model cost map endpoint""" # Mock non-admin user mock_auth = MagicMock() mock_auth.user_role = "user" # Non-admin role app.dependency_overrides[user_api_key_auth] = lambda: mock_auth - + response = client_with_auth.get("/get/litellm_model_cost_map") - + assert response.status_code == 403 data = response.json() assert "Access denied" in data["detail"] assert "Admin role required" in data["detail"] - + def test_reload_model_cost_map_error_handling(self, client_with_auth): """Test error handling in the reload endpoint""" - with patch('litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map') as mock_get_map: + with patch( + "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map" + ) as mock_get_map: mock_get_map.side_effect = Exception("Network error") - + # Mock the database connection - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) - + response = client_with_auth.post("/reload/model_cost_map") - - assert response.status_code == 500 # The new implementation immediately reloads and fails on error + + assert ( + response.status_code == 500 + ) # The new implementation immediately reloads and fails on error data = response.json() assert "Failed to reload model cost map" in data["detail"] def test_schedule_model_cost_map_reload_admin_access(self, client_with_auth): """Test that admin users can schedule periodic reload""" - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # Mock database upsert mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) - + response = client_with_auth.post("/schedule/model_cost_map_reload?hours=6") - + assert response.status_code == 200 data = response.json() assert data["status"] == "success" @@ -1312,9 +1328,9 @@ class TestPriceDataReloadAPI: mock_auth = MagicMock() mock_auth.user_role = "user" # Non-admin role app.dependency_overrides[user_api_key_auth] = lambda: mock_auth - + response = client_with_auth.post("/schedule/model_cost_map_reload?hours=6") - + assert response.status_code == 403 data = response.json() assert "Access denied" in data["detail"] @@ -1323,19 +1339,19 @@ class TestPriceDataReloadAPI: def test_schedule_model_cost_map_reload_invalid_hours(self, client_with_auth): """Test that invalid hours parameter is rejected""" response = client_with_auth.post("/schedule/model_cost_map_reload?hours=0") - + assert response.status_code == 400 data = response.json() assert "Hours must be greater than 0" in data["detail"] def test_cancel_model_cost_map_reload_admin_access(self, client_with_auth): """Test that admin users can cancel periodic reload""" - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # Mock database delete mock_prisma.db.litellm_config.delete = AsyncMock(return_value=None) - + response = client_with_auth.delete("/schedule/model_cost_map_reload") - + assert response.status_code == 200 data = response.json() assert data["status"] == "success" @@ -1348,9 +1364,9 @@ class TestPriceDataReloadAPI: mock_auth = MagicMock() mock_auth.user_role = "user" # Non-admin role app.dependency_overrides[user_api_key_auth] = lambda: mock_auth - + response = client_with_auth.delete("/schedule/model_cost_map_reload") - + assert response.status_code == 403 data = response.json() assert "Access denied" in data["detail"] @@ -1358,24 +1374,28 @@ class TestPriceDataReloadAPI: def test_get_model_cost_map_reload_status_admin_access(self, client_with_auth): """Test that admin users can get reload status""" - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # Mock database config record mock_config = MagicMock() - mock_config.param_value = { - "interval_hours": 6, - "force_reload": False - } - mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config) - + mock_config.param_value = {"interval_hours": 6, "force_reload": False} + mock_prisma.db.litellm_config.find_unique = AsyncMock( + return_value=mock_config + ) + # Mock the last reload time and current time - with patch('litellm.proxy.proxy_server.last_model_cost_map_reload', "2024-01-01T06:00:00"): - with patch('litellm.proxy.proxy_server.datetime') as mock_datetime: + with patch( + "litellm.proxy.proxy_server.last_model_cost_map_reload", + "2024-01-01T06:00:00", + ): + with patch("litellm.proxy.proxy_server.datetime") as mock_datetime: # Mock current time to be 1 hour after last reload mock_datetime.utcnow.return_value = datetime(2024, 1, 1, 7, 0, 0) mock_datetime.fromisoformat = datetime.fromisoformat - - response = client_with_auth.get("/schedule/model_cost_map_reload/status") - + + response = client_with_auth.get( + "/schedule/model_cost_map_reload/status" + ) + assert response.status_code == 200 data = response.json() assert data["scheduled"] == True @@ -1389,9 +1409,9 @@ class TestPriceDataReloadAPI: mock_auth = MagicMock() mock_auth.user_role = "user" # Non-admin role app.dependency_overrides[user_api_key_auth] = lambda: mock_auth - + response = client_with_auth.get("/schedule/model_cost_map_reload/status") - + assert response.status_code == 403 data = response.json() assert "Access denied" in data["detail"] @@ -1399,11 +1419,11 @@ class TestPriceDataReloadAPI: def test_get_model_cost_map_reload_status_no_config(self, client_with_auth): """Test that status returns not scheduled when no config exists""" - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) - + response = client_with_auth.get("/schedule/model_cost_map_reload/status") - + assert response.status_code == 200 data = response.json() assert data["scheduled"] == False @@ -1413,17 +1433,16 @@ class TestPriceDataReloadAPI: def test_get_model_cost_map_reload_status_no_interval(self, client_with_auth): """Test that status returns not scheduled when no interval is configured""" - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # Mock config with no interval mock_config = MagicMock() - mock_config.param_value = { - "interval_hours": None, - "force_reload": False - } - mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config) - + mock_config.param_value = {"interval_hours": None, "force_reload": False} + mock_prisma.db.litellm_config.find_unique = AsyncMock( + return_value=mock_config + ) + response = client_with_auth.get("/schedule/model_cost_map_reload/status") - + assert response.status_code == 200 data = response.json() assert data["scheduled"] == False @@ -1434,107 +1453,109 @@ class TestPriceDataReloadAPI: class TestPriceDataReloadIntegration: """Integration tests for the complete price data reload feature""" - + @pytest.fixture def client_with_auth(self): """Create a test client with authentication""" - from litellm.proxy.proxy_server import cleanup_router_config_variables from litellm.proxy._types import LitellmUserRoles - + from litellm.proxy.proxy_server import cleanup_router_config_variables + cleanup_router_config_variables() filepath = os.path.dirname(os.path.abspath(__file__)) config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml" asyncio.run(initialize(config=config_fp, debug=True)) - + # Mock admin user authentication mock_auth = MagicMock() mock_auth.user_role = LitellmUserRoles.PROXY_ADMIN app.dependency_overrides[user_api_key_auth] = lambda: mock_auth - + return TestClient(app) - + def test_complete_reload_flow(self, client_with_auth): """Test the complete reload flow from API to model cost update""" # Mock the model cost map mock_cost_map = { "gpt-3.5-turbo": { "input_cost_per_token": 0.001, - "output_cost_per_token": 0.002 + "output_cost_per_token": 0.002, }, - "gpt-4": { - "input_cost_per_token": 0.03, - "output_cost_per_token": 0.06 - } + "gpt-4": {"input_cost_per_token": 0.03, "output_cost_per_token": 0.06}, } - - with patch('litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map') as mock_get_map: + + with patch( + "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map" + ) as mock_get_map: mock_get_map.return_value = mock_cost_map - + # Mock the database connection - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) - + # Test reload endpoint response = client_with_auth.post("/reload/model_cost_map") assert response.status_code == 200 - + # Test get endpoint response = client_with_auth.get("/get/litellm_model_cost_map") assert response.status_code == 200 - + def test_distributed_reload_check_function(self): """Test the _check_and_reload_model_cost_map function""" from litellm.proxy.proxy_server import ProxyConfig - + proxy_config = ProxyConfig() - + # Mock prisma client mock_prisma = MagicMock() - + # Test case 1: No config in database mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) - + # Should return early without reloading asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) - + # Test case 2: Config with interval but not time to reload mock_config = MagicMock() - mock_config.param_value = { - "interval_hours": 6, - "force_reload": False - } + mock_config.param_value = {"interval_hours": 6, "force_reload": False} mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config) - + # Mock current time and last reload time - with patch('litellm.proxy.proxy_server.last_model_cost_map_reload', "2024-01-01T06:00:00"): - with patch('litellm.proxy.proxy_server.datetime') as mock_datetime: - mock_datetime.utcnow.return_value = datetime(2024, 1, 1, 7, 0, 0) # 1 hour later - + with patch( + "litellm.proxy.proxy_server.last_model_cost_map_reload", + "2024-01-01T06:00:00", + ): + with patch("litellm.proxy.proxy_server.datetime") as mock_datetime: + mock_datetime.utcnow.return_value = datetime( + 2024, 1, 1, 7, 0, 0 + ) # 1 hour later + # Should not reload (only 1 hour passed, need 6) asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) - + # Test case 3: Config with force reload - mock_config.param_value = { - "interval_hours": 6, - "force_reload": True - } + mock_config.param_value = {"interval_hours": 6, "force_reload": True} mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config) mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) - - with patch('litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map') as mock_get_map: - mock_get_map.return_value = {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}} - + + with patch( + "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map" + ) as mock_get_map: + mock_get_map.return_value = { + "gpt-3.5-turbo": {"input_cost_per_token": 0.001} + } + # Should reload due to force flag asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) - + # Verify force_reload was reset to False mock_prisma.db.litellm_config.upsert.assert_called() call_args = mock_prisma.db.litellm_config.upsert.call_args # The param_value is now a JSON string, so we need to parse it - param_value_json = call_args[1]['data']['update']['param_value'] + param_value_json = call_args[1]["data"]["update"]["param_value"] param_value_dict = json.loads(param_value_json) - assert param_value_dict['force_reload'] == False - + assert param_value_dict["force_reload"] == False + def test_config_file_parsing(self): """Test parsing of config file with reload settings""" config_content = """ @@ -1550,81 +1571,315 @@ model_list: litellm_params: model: gpt-4 """ - + # Parse the config config = yaml.safe_load(config_content) - + # Verify the reload setting is present assert "general_settings" in config assert "model_cost_map_reload_interval" in config["general_settings"] assert config["general_settings"]["model_cost_map_reload_interval"] == 21600 - + # Verify models are present assert "model_list" in config assert len(config["model_list"]) == 2 + def test_database_config_storage(self): """Test that configuration is properly stored in database""" # Mock prisma client mock_prisma = MagicMock() - + # Test the database upsert call that would be made by the schedule endpoint mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) - + # Simulate the database call that the schedule endpoint would make - asyncio.run(mock_prisma.db.litellm_config.upsert( - where={"param_name": "model_cost_map_reload_config"}, - data={ - "create": { - "param_name": "model_cost_map_reload_config", - "param_value": { - "interval_hours": 6, - "force_reload": False - } + asyncio.run( + mock_prisma.db.litellm_config.upsert( + where={"param_name": "model_cost_map_reload_config"}, + data={ + "create": { + "param_name": "model_cost_map_reload_config", + "param_value": {"interval_hours": 6, "force_reload": False}, + }, + "update": { + "param_value": {"interval_hours": 6, "force_reload": False} + }, }, - "update": { - "param_value": { - "interval_hours": 6, - "force_reload": False - } - } - } - )) - + ) + ) + # Verify database upsert was called with correct data mock_prisma.db.litellm_config.upsert.assert_called_once() call_args = mock_prisma.db.litellm_config.upsert.call_args - assert call_args[1]['where']['param_name'] == "model_cost_map_reload_config" - assert call_args[1]['data']['create']['param_value']['interval_hours'] == 6 - assert call_args[1]['data']['create']['param_value']['force_reload'] == False + assert call_args[1]["where"]["param_name"] == "model_cost_map_reload_config" + assert call_args[1]["data"]["create"]["param_value"]["interval_hours"] == 6 + assert call_args[1]["data"]["create"]["param_value"]["force_reload"] == False def test_manual_reload_force_flag(self): """Test that manual reload sets force flag correctly""" # Mock prisma client mock_prisma = MagicMock() - + # Test the database upsert call that would be made by the manual reload endpoint mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) - + # Simulate the database call that the manual reload endpoint would make - asyncio.run(mock_prisma.db.litellm_config.upsert( - where={"param_name": "model_cost_map_reload_config"}, - data={ - "create": { - "param_name": "model_cost_map_reload_config", - "param_value": { - "interval_hours": None, - "force_reload": True - } + asyncio.run( + mock_prisma.db.litellm_config.upsert( + where={"param_name": "model_cost_map_reload_config"}, + data={ + "create": { + "param_name": "model_cost_map_reload_config", + "param_value": {"interval_hours": None, "force_reload": True}, + }, + "update": {"param_value": {"force_reload": True}}, }, - "update": { - "param_value": { - "force_reload": True - } - } - } - )) - + ) + ) + # Verify force_reload flag was set mock_prisma.db.litellm_config.upsert.assert_called_once() call_args = mock_prisma.db.litellm_config.upsert.call_args - assert call_args[1]['data']['update']['param_value']['force_reload'] == True + assert call_args[1]["data"]["update"]["param_value"]["force_reload"] == True + + +@pytest.mark.asyncio +async def test_add_router_settings_from_db_config_merge_logic(): + """ + Test the _add_router_settings_from_db_config method's merge logic. + + This tests how router settings from config file and database are combined, + including scenarios where nested dictionaries should be properly merged. + """ + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy.proxy_server import ProxyConfig + + # Create ProxyConfig instance + proxy_config = ProxyConfig() + + # Mock router + mock_router = MagicMock() + mock_router.update_settings = MagicMock() + + # Test Case 1: Both config and DB settings exist - should merge them + config_data = { + "router_settings": { + "routing_strategy": "usage-based-routing", + "model_group_alias": {"gpt-4": "openai-gpt-4"}, + "enable_pre_call_checks": True, + "timeout": 30, + "nested_config": {"setting1": "config_value1", "setting2": "config_value2"}, + } + } + + # Mock database config record + mock_db_config = MagicMock() + mock_db_config.param_value = { + "routing_strategy": "least-busy", # This should override config value + "retry_delay": 2, # This is new, should be added + "nested_config": { + "setting2": "db_value2", # This should override config value + "setting3": "db_value3", # This is new, should be added + }, + } + + # Mock prisma client + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_config.find_first = AsyncMock( + return_value=mock_db_config + ) + + # Call the method under test + await proxy_config._add_router_settings_from_db_config( + config_data=config_data, + llm_router=mock_router, + prisma_client=mock_prisma_client, + ) + + # Verify find_first was called with correct parameters + mock_prisma_client.db.litellm_config.find_first.assert_called_once_with( + where={"param_name": "router_settings"} + ) + + # Verify update_settings was called + mock_router.update_settings.assert_called_once() + + # Get the actual settings passed to update_settings + call_args = mock_router.update_settings.call_args + combined_settings = call_args[1] # kwargs + + # Verify the merge results + # DB values should override config values + assert combined_settings["routing_strategy"] == "least-busy" + + # Config-only values should be preserved + assert combined_settings["model_group_alias"] == {"gpt-4": "openai-gpt-4"} + assert combined_settings["enable_pre_call_checks"] == True + assert combined_settings["timeout"] == 30 + + # DB-only values should be added + assert combined_settings["retry_delay"] == 2 + + # Nested dictionaries should be merged (but this is shallow merge) + # The entire nested_config dict gets replaced by DB value + expected_nested = {"setting2": "db_value2", "setting3": "db_value3"} + assert combined_settings["nested_config"] == expected_nested + + +@pytest.mark.asyncio +async def test_add_router_settings_from_db_config_edge_cases(): + """ + Test edge cases for _add_router_settings_from_db_config method. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + mock_router = MagicMock() + mock_router.update_settings = MagicMock() + + # Test Case 1: No router provided + await proxy_config._add_router_settings_from_db_config( + config_data={"router_settings": {"test": "value"}}, + llm_router=None, + prisma_client=MagicMock(), + ) + # Should not call anything when router is None + mock_router.update_settings.assert_not_called() + + # Test Case 2: No prisma client provided + await proxy_config._add_router_settings_from_db_config( + config_data={"router_settings": {"test": "value"}}, + llm_router=mock_router, + prisma_client=None, + ) + # Should not call anything when prisma_client is None + mock_router.update_settings.assert_not_called() + + # Test Case 3: DB returns None (no router_settings in DB) + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) + + config_data = {"router_settings": {"routing_strategy": "usage-based"}} + + await proxy_config._add_router_settings_from_db_config( + config_data=config_data, + llm_router=mock_router, + prisma_client=mock_prisma_client, + ) + + # Should use only config settings + mock_router.update_settings.assert_called_once_with(routing_strategy="usage-based") + mock_router.reset_mock() + + # Test Case 4: Config has no router_settings + mock_db_config = MagicMock() + mock_db_config.param_value = {"db_setting": "db_value"} + mock_prisma_client.db.litellm_config.find_first = AsyncMock( + return_value=mock_db_config + ) + + await proxy_config._add_router_settings_from_db_config( + config_data={}, # No router_settings in config + llm_router=mock_router, + prisma_client=mock_prisma_client, + ) + + # Should use only DB settings + mock_router.update_settings.assert_called_once_with(db_setting="db_value") + mock_router.reset_mock() + + # Test Case 5: Both config and DB router_settings are None/empty + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) + + await proxy_config._add_router_settings_from_db_config( + config_data={}, llm_router=mock_router, prisma_client=mock_prisma_client + ) + + # Should not call update_settings when no settings exist + mock_router.update_settings.assert_not_called() + + # Test Case 6: DB config exists but param_value is not a dict + mock_db_config_invalid = MagicMock() + mock_db_config_invalid.param_value = "not_a_dict" + mock_prisma_client.db.litellm_config.find_first = AsyncMock( + return_value=mock_db_config_invalid + ) + + config_data = {"router_settings": {"config_setting": "config_value"}} + + await proxy_config._add_router_settings_from_db_config( + config_data=config_data, + llm_router=mock_router, + prisma_client=mock_prisma_client, + ) + + # Should use only config settings when DB param_value is invalid + mock_router.update_settings.assert_called_once_with(config_setting="config_value") + + +@pytest.mark.asyncio +async def test_add_router_settings_shallow_merge_behavior(): + """ + Test that the merge behavior is shallow (nested dicts get replaced, not merged). + This documents the current behavior using _update_dictionary. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + mock_router = MagicMock() + mock_router.update_settings = MagicMock() + + # Config with nested dictionary + config_data = { + "router_settings": { + "nested_setting": { + "key1": "config_value1", + "key2": "config_value2", + "key3": "config_value3", + }, + "top_level": "config_top", + } + } + + # DB config that partially overlaps the nested dictionary + mock_db_config = MagicMock() + mock_db_config.param_value = { + "nested_setting": { + "key2": "db_value2", # Override existing key + "key4": "db_value4", # Add new key + # Note: key1 and key3 from config will be lost due to shallow merge + }, + "top_level": "db_top", # Override top level + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_config.find_first = AsyncMock( + return_value=mock_db_config + ) + + await proxy_config._add_router_settings_from_db_config( + config_data=config_data, + llm_router=mock_router, + prisma_client=mock_prisma_client, + ) + + # Get the merged settings + call_args = mock_router.update_settings.call_args + merged_settings = call_args[1] + + # Verify shallow merge behavior: + # The entire nested_setting dict from config is replaced by the DB version + expected_nested = { + "key1": "config_value1", + "key3": "config_value3", + "key2": "db_value2", + "key4": "db_value4", + } + + assert merged_settings["nested_setting"] == expected_nested + assert merged_settings["top_level"] == "db_top"