mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy_server.py): fix merge config + db router settings logic
This commit is contained in:
parent
52cbedf044
commit
e5c3e006c2
3 changed files with 457 additions and 170 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue