mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
[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:
parent
e87e50328e
commit
f4318bccd3
3 changed files with 161 additions and 21 deletions
|
|
@ -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 ###
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue