diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index 83d7e173431..467565d748a 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -17,13 +17,13 @@ except ImportError: # List of all available hooks that can be enabled PROXY_HOOKS = { "max_budget_limiter": _PROXY_MaxBudgetLimiter, - "parallel_request_limiter": _PROXY_MaxParallelRequestsHandler, + "parallel_request_limiter": _PROXY_MaxParallelRequestsHandler_v3, "cache_control_check": _PROXY_CacheControlCheck, } ## FEATURE FLAG HOOKS ## -if os.getenv("EXPERIMENTAL_MULTI_INSTANCE_RATE_LIMITING", "false").lower() == "true": - PROXY_HOOKS["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3 +if os.getenv("LEGACY_MULTI_INSTANCE_RATE_LIMITING", "false").lower() == "true": + PROXY_HOOKS["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler ### update PROXY_HOOKS with ENTERPRISE_PROXY_HOOKS ### diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b3840761d2a..3a1a0bef176 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -266,7 +266,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if current_limit is None or rate_limit_type is None: continue - if counter_value is not None and int(counter_value) + 1 > current_limit: + if counter_value is not None and int(counter_value) > current_limit: overall_code = "OVER_LIMIT" item_code = "OVER_LIMIT" diff --git a/tests/local_testing/test_pass_through_endpoints.py b/tests/local_testing/test_pass_through_endpoints.py index 118a1f64f41..6cc6a66007f 100644 --- a/tests/local_testing/test_pass_through_endpoints.py +++ b/tests/local_testing/test_pass_through_endpoints.py @@ -1,5 +1,7 @@ import os import sys +import uuid +from functools import partial from typing import Optional import pytest @@ -146,24 +148,36 @@ async def test_pass_through_endpoint_rerank(client): @pytest.mark.parametrize( - "auth, rpm_limit, expected_error_code", - [(True, 0, 429), (True, 1, 200), (False, 0, 200)], + "auth, rpm_limit, requests_to_make, expected_status_codes, num_users", + [ + # Single user tests + (True, 0, 1, [429], 1), + (True, 1, 1, [200], 1), + (True, 1, 2, [200, 429], 1), + (True, 2, 4, [200, 200, 429, 429], 1), + (True, 3, 4, [200, 200, 200, 429], 1), + (True, 4, 4, [200, 200, 200, 200], 1), + (False, 0, 1, [200], 1), + (False, 0, 4, [200, 200, 200, 200], 1), + # Multiple user tests (same parameters as single user) + (True, 0, 1, [429], 2), + (True, 1, 1, [200], 2), + (True, 1, 2, [200, 429], 2), + (True, 2, 4, [200, 200, 429, 429], 2), + (True, 3, 4, [200, 200, 200, 429], 2), + (True, 4, 4, [200, 200, 200, 200], 2), + (False, 0, 1, [200], 2), + (False, 0, 4, [200, 200, 200, 200], 2), + ], ) @pytest.mark.asyncio async def test_pass_through_endpoint_rpm_limit( - client, auth, expected_error_code, rpm_limit + client, auth, rpm_limit, requests_to_make, expected_status_codes, num_users ): import litellm from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache - mock_api_key = "sk-my-test-key" - cache_value = UserAPIKeyAuth(token=hash_token(mock_api_key), rpm_limit=rpm_limit) - - _cohere_api_key = os.environ.get("COHERE_API_KEY") - - user_api_key_cache.set_cache(key=hash_token(mock_api_key), value=cache_value) - proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) proxy_logging_obj._init_litellm_callbacks() @@ -173,6 +187,7 @@ async def test_pass_through_endpoint_rpm_limit( setattr(litellm.proxy.proxy_server, "proxy_logging_obj", proxy_logging_obj) # Define a pass-through endpoint + _cohere_api_key = os.environ.get("COHERE_API_KEY") pass_through_endpoints = [ { "path": "/v1/rerank", @@ -190,6 +205,13 @@ async def test_pass_through_endpoint_rpm_limit( general_settings.update({"pass_through_endpoints": pass_through_endpoints}) setattr(litellm.proxy.proxy_server, "general_settings", general_settings) + # Setup API keys and cache + mock_api_keys = [f"sk-test-{uuid.uuid4().hex}" for _ in range(num_users)] + + for mock_api_key in mock_api_keys: + cache_value = UserAPIKeyAuth(token=hash_token(mock_api_key), rpm_limit=rpm_limit) + user_api_key_cache.set_cache(key=hash_token(mock_api_key), value=cache_value) + _json_data = { "model": "rerank-english-v3.0", "query": "What is the capital of the United States?", @@ -200,16 +222,134 @@ async def test_pass_through_endpoint_rpm_limit( } # Make a request to the pass-through endpoint - response = client.post( - "/v1/rerank", - json=_json_data, - headers={"Authorization": "Bearer {}".format(mock_api_key)}, - ) + tasks = [] + for mock_api_key in mock_api_keys: + for _ in range(requests_to_make): + task = asyncio.get_running_loop().run_in_executor( + None, + partial( + client.post, + "/v1/rerank", + json=_json_data, + headers={"Authorization": "Bearer {}".format(mock_api_key)}, + ), + ) + tasks.append(task) + + responses = await asyncio.gather(*tasks) + + if num_users == 1: + status_codes = sorted([response.status_code for response in responses]) + + assert status_codes == sorted(expected_status_codes) + else: + first_user_responses = responses[requests_to_make:] + second_user_responses = responses[:requests_to_make] + + first_user_status_codes = sorted([response.status_code for response in first_user_responses]) + second_user_status_codes = sorted([response.status_code for response in second_user_responses]) + + expected_status_codes.sort() + assert first_user_status_codes == expected_status_codes + assert second_user_status_codes == expected_status_codes print("JSON response: ", _json_data) - # Assert the response - assert response.status_code == expected_error_code + +@pytest.mark.parametrize( + "auth, rpm_limit, requests_to_make, expected_status_codes", + [ + # Multiple user tests (same parameters as single user) + (True, 0, 1, [429]), + (True, 1, 1, [200]), + (True, 1, 2, [200, 429]), + (True, 2, 4, [200, 200, 429, 429]), + (True, 3, 4, [200, 200, 200, 429]), + (True, 4, 4, [200, 200, 200, 200]), + (False, 0, 1, [200]), + (False, 0, 4, [200, 200, 200, 200]), + ], +) +@pytest.mark.asyncio +async def test_pass_through_endpoint_sequential_rpm_limit( + client, auth, rpm_limit, requests_to_make, expected_status_codes +): + import litellm + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache + + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + proxy_logging_obj._init_litellm_callbacks() + + setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + setattr(litellm.proxy.proxy_server, "prisma_client", "FAKE-VAR") + setattr(litellm.proxy.proxy_server, "proxy_logging_obj", proxy_logging_obj) + + # Define a pass-through endpoint + _cohere_api_key = os.environ.get("COHERE_API_KEY") + pass_through_endpoints = [ + { + "path": "/v1/rerank", + "target": "https://api.cohere.com/v1/rerank", + "auth": auth, + "headers": {"Authorization": f"bearer {_cohere_api_key}"}, + } + ] + + # Initialize the pass-through endpoint + await initialize_pass_through_endpoints(pass_through_endpoints) + general_settings: Optional[dict] = ( + getattr(litellm.proxy.proxy_server, "general_settings", {}) or {} + ) + general_settings.update({"pass_through_endpoints": pass_through_endpoints}) + setattr(litellm.proxy.proxy_server, "general_settings", general_settings) + + # Setup API keys and cache + mock_api_keys = [f"sk-test-{uuid.uuid4().hex}" for _ in range(2)] + + for mock_api_key in mock_api_keys: + cache_value = UserAPIKeyAuth(token=hash_token(mock_api_key), rpm_limit=rpm_limit) + user_api_key_cache.set_cache(key=hash_token(mock_api_key), value=cache_value) + + _json_data = { + "model": "rerank-english-v3.0", + "query": "What is the capital of the United States?", + "top_n": 3, + "documents": [ + "Carson City is the capital city of the American state of Nevada." + ], + } + + # Make a request to the pass-through endpoint + first_user_responses = [] + second_user_responses = [] + for _ in range(requests_to_make): + requests = [] + for mock_api_key in mock_api_keys: + task = asyncio.get_running_loop().run_in_executor( + None, + partial( + client.post, + "/v1/rerank", + json=_json_data, + headers={"Authorization": "Bearer {}".format(mock_api_key)}, + ), + ) + requests.append(task) + + first_user_response, second_user_response = await asyncio.gather(*requests) + first_user_responses.append(first_user_response) + second_user_responses.append(second_user_response) + + first_user_status_codes = sorted([response.status_code for response in first_user_responses]) + second_user_status_codes = sorted([response.status_code for response in second_user_responses]) + + expected_status_codes.sort() + assert first_user_status_codes == expected_status_codes + assert second_user_status_codes == expected_status_codes + + print("JSON response: ", _json_data) @pytest.mark.parametrize(