fix(proxy_server.py): fix merge config + db router settings logic

This commit is contained in:
Krrish Dholakia 2025-08-20 16:33:31 -07:00
parent 52cbedf044
commit e5c3e006c2
3 changed files with 457 additions and 170 deletions

View file

@ -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

View file

@ -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()

View file

@ -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"