[Performance] Use _PROXY_MaxParallelRequestsHandler_v3 by default again (#14450)

* Use _PROXY_MaxParallelRequestsHandler_v3 by default (#14352)

(cherry picked from commit f3fa45cf8fbd5f5cce2f45a7312776d5005fb08e)
(cherry picked from commit 5b680bb4a3)

* Use random api_key for parallel requests test

* Fix off-by-one error in parallel request rate limit

The rate limiter was incorrectly rejecting requests when the limit was met, but not exceeded. The check in `is_cache_list_over_limit` was `int(counter_value) + 1 > current_limit`, which caused the first request to be rejected if the limit was 1.

This commit removes the `+ 1`, changing the logic to `int(counter_value) > current_limit`. The check now correctly allows requests up to the specified parallel limit.

* Test actual parallel requests

* Ensure rate limiting works correctly for multiple users

* Add sequential rate-limit test

* Revert random key usage
This commit is contained in:
Arseny Boykov 2025-09-13 02:33:55 +02:00 • committed by GitHub
parent e87e50328e
commit f4318bccd3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 161 additions and 21 deletions

View file

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

View file

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

View file

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