diff --git a/tests/litellm/proxy/test_litellm_pre_call_utils.py b/tests/litellm/proxy/test_litellm_pre_call_utils.py index b671f71506..94f2c512ea 100644 --- a/tests/litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/litellm/proxy/test_litellm_pre_call_utils.py @@ -61,3 +61,45 @@ def test_get_enforced_params_for_service_account_settings(): user_api_key_dict=regular_token, ) assert result == ["user"] + + +@pytest.mark.parametrize( + "general_settings, user_api_key_dict, expected_enforced_params", + [ + ( + {"enforced_params": ["param1", "param2"]}, + UserAPIKeyAuth( + api_key="test_api_key", user_id="test_user_id", org_id="test_org_id" + ), + ["param1", "param2"], + ), + ( + {"service_account_settings": {"enforced_params": ["param1", "param2"]}}, + UserAPIKeyAuth( + api_key="test_api_key", + user_id="test_user_id", + org_id="test_org_id", + metadata={"service_account_id": "test_service_account_id"}, + ), + ["param1", "param2"], + ), + ( + {"service_account_settings": {"enforced_params": ["param1", "param2"]}}, + UserAPIKeyAuth( + api_key="test_api_key", + metadata={ + "enforced_params": ["param3", "param4"], + "service_account_id": "test_service_account_id", + }, + ), + ["param1", "param2", "param3", "param4"], + ), + ], +) +def test_get_enforced_params( + general_settings, user_api_key_dict, expected_enforced_params +): + from litellm.proxy.litellm_pre_call_utils import _get_enforced_params + + enforced_params = _get_enforced_params(general_settings, user_api_key_dict) + assert enforced_params == expected_enforced_params diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 9fb00c094f..b28948094e 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -769,42 +769,6 @@ async def test_add_litellm_data_to_request_duplicate_tags( ), f"Expected {expected_tags}, got {result['metadata']['tags']}" -@pytest.mark.parametrize( - "general_settings, user_api_key_dict, expected_enforced_params", - [ - ( - {"enforced_params": ["param1", "param2"]}, - UserAPIKeyAuth( - api_key="test_api_key", user_id="test_user_id", org_id="test_org_id" - ), - ["param1", "param2"], - ), - ( - {"service_account_settings": {"enforced_params": ["param1", "param2"]}}, - UserAPIKeyAuth( - api_key="test_api_key", user_id="test_user_id", org_id="test_org_id" - ), - ["param1", "param2"], - ), - ( - {"service_account_settings": {"enforced_params": ["param1", "param2"]}}, - UserAPIKeyAuth( - api_key="test_api_key", - metadata={"enforced_params": ["param3", "param4"]}, - ), - ["param1", "param2", "param3", "param4"], - ), - ], -) -def test_get_enforced_params( - general_settings, user_api_key_dict, expected_enforced_params -): - from litellm.proxy.litellm_pre_call_utils import _get_enforced_params - - enforced_params = _get_enforced_params(general_settings, user_api_key_dict) - assert enforced_params == expected_enforced_params - - @pytest.mark.parametrize( "general_settings, user_api_key_dict, request_body, expected_error", [