mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: enforce project-level model-specific rate limits in parallel_request_limiter_v3
Project-level model rpm/tpm limits stored in project_metadata were never
checked during rate limit enforcement — only model-level limits applied.
Adds _add_project_model_rate_limit_descriptor_from_metadata() to the v3
limiter (mirrors the existing team metadata path) and calls it in
async_pre_call_hook, creating a model_per_project descriptor keyed as
"{project_id}:{model}" with the project's configured limits.
Also extends get_model_rate_limit_from_metadata's Literal to accept
"project_metadata" and adds get_project_model_rpm/tpm_limit helpers.
Fixes: LIT-2317
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
850fe595ac
commit
6fd49f1da1
3 changed files with 171 additions and 1 deletions
|
|
@ -663,7 +663,9 @@ def get_key_model_tpm_limit(
|
|||
|
||||
def get_model_rate_limit_from_metadata(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
metadata_accessor_key: Literal["team_metadata", "organization_metadata"],
|
||||
metadata_accessor_key: Literal[
|
||||
"team_metadata", "organization_metadata", "project_metadata"
|
||||
],
|
||||
rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"],
|
||||
) -> Optional[Dict[str, int]]:
|
||||
if getattr(user_api_key_dict, metadata_accessor_key):
|
||||
|
|
@ -687,6 +689,22 @@ def get_team_model_tpm_limit(
|
|||
return None
|
||||
|
||||
|
||||
def get_project_model_rpm_limit(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Optional[Dict[str, int]]:
|
||||
if user_api_key_dict.project_metadata:
|
||||
return user_api_key_dict.project_metadata.get("model_rpm_limit")
|
||||
return None
|
||||
|
||||
|
||||
def get_project_model_tpm_limit(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Optional[Dict[str, int]]:
|
||||
if user_api_key_dict.project_metadata:
|
||||
return user_api_key_dict.project_metadata.get("model_tpm_limit")
|
||||
return None
|
||||
|
||||
|
||||
def is_pass_through_provider_route(route: str) -> bool:
|
||||
PROVIDER_SPECIFIC_PASS_THROUGH_ROUTES = [
|
||||
"vertex-ai",
|
||||
|
|
|
|||
|
|
@ -1175,6 +1175,59 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
)
|
||||
|
||||
def _add_project_model_rate_limit_descriptor_from_metadata(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
requested_model: Optional[str],
|
||||
descriptors: List[RateLimitDescriptor],
|
||||
) -> None:
|
||||
"""Add project model rate limit descriptor from project_metadata if applicable."""
|
||||
if (
|
||||
get_model_rate_limit_from_metadata(
|
||||
user_api_key_dict, "project_metadata", "model_rpm_limit"
|
||||
)
|
||||
is not None
|
||||
or get_model_rate_limit_from_metadata(
|
||||
user_api_key_dict, "project_metadata", "model_tpm_limit"
|
||||
)
|
||||
is not None
|
||||
):
|
||||
_tpm_limit_for_project_model = (
|
||||
get_model_rate_limit_from_metadata(
|
||||
user_api_key_dict, "project_metadata", "model_tpm_limit"
|
||||
)
|
||||
or {}
|
||||
)
|
||||
_rpm_limit_for_project_model = (
|
||||
get_model_rate_limit_from_metadata(
|
||||
user_api_key_dict, "project_metadata", "model_rpm_limit"
|
||||
)
|
||||
or {}
|
||||
)
|
||||
should_check_rate_limit = (
|
||||
requested_model in _tpm_limit_for_project_model
|
||||
or requested_model in _rpm_limit_for_project_model
|
||||
)
|
||||
|
||||
if should_check_rate_limit and requested_model is not None:
|
||||
model_specific_tpm_limit = _tpm_limit_for_project_model.get(
|
||||
requested_model
|
||||
)
|
||||
model_specific_rpm_limit = _rpm_limit_for_project_model.get(
|
||||
requested_model
|
||||
)
|
||||
descriptors.append(
|
||||
RateLimitDescriptor(
|
||||
key="model_per_project",
|
||||
value=f"{user_api_key_dict.project_id}:{requested_model}",
|
||||
rate_limit={
|
||||
"requests_per_unit": model_specific_rpm_limit,
|
||||
"tokens_per_unit": model_specific_tpm_limit,
|
||||
"window_size": self.window_size,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
def _handle_rate_limit_error(
|
||||
self,
|
||||
response: RateLimitResponse,
|
||||
|
|
@ -1286,6 +1339,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
descriptors=descriptors,
|
||||
)
|
||||
|
||||
# Project Level Rate Limits
|
||||
self._add_project_model_rate_limit_descriptor_from_metadata(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
requested_model=requested_model,
|
||||
descriptors=descriptors,
|
||||
)
|
||||
|
||||
# Org Level Rate Limits
|
||||
descriptors.extend(
|
||||
self.create_organization_rate_limit_descriptor(
|
||||
|
|
|
|||
|
|
@ -2590,3 +2590,95 @@ class TestGetTotalTokensFromUsageCacheExclusion:
|
|||
"""Should handle None usage gracefully."""
|
||||
result = handler._get_total_tokens_from_usage(None, "total")
|
||||
assert result == 0, f"Expected 0 for None usage, got {result}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_project_model_rate_limits_enforced_v3():
|
||||
"""
|
||||
Regression test: project-level model-specific rate limits must be enforced.
|
||||
|
||||
Bug: When a key belongs to a project that has model_rpm_limit/model_tpm_limit
|
||||
in project_metadata, those limits were never checked — only model-level limits
|
||||
were applied. This test verifies the fix.
|
||||
"""
|
||||
_api_key = hash_token("sk-project-test")
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
|
||||
captured_descriptors = []
|
||||
|
||||
async def mock_should_rate_limit(descriptors, **kwargs):
|
||||
captured_descriptors.extend(descriptors)
|
||||
return {"overall_code": "OK", "statuses": []}
|
||||
|
||||
parallel_request_handler.should_rate_limit = mock_should_rate_limit
|
||||
|
||||
# Key with project_metadata containing model-specific rate limits
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key=_api_key,
|
||||
project_id="proj-abc123",
|
||||
project_metadata={
|
||||
"model_rpm_limit": {"gpt-4": 5},
|
||||
"model_tpm_limit": {"gpt-4": 1000},
|
||||
},
|
||||
)
|
||||
|
||||
await parallel_request_handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=local_cache,
|
||||
data={"model": "gpt-4"},
|
||||
call_type="",
|
||||
)
|
||||
|
||||
descriptor_keys = [d["key"] for d in captured_descriptors]
|
||||
assert (
|
||||
"model_per_project" in descriptor_keys
|
||||
), f"Expected model_per_project descriptor, got: {descriptor_keys}"
|
||||
|
||||
model_per_project = next(
|
||||
d for d in captured_descriptors if d["key"] == "model_per_project"
|
||||
)
|
||||
assert model_per_project["value"] == "proj-abc123:gpt-4"
|
||||
assert model_per_project["rate_limit"]["requests_per_unit"] == 5
|
||||
assert model_per_project["rate_limit"]["tokens_per_unit"] == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_project_model_rate_limits_not_triggered_for_other_model_v3():
|
||||
"""Project model limits should not trigger for a model not in project_metadata."""
|
||||
_api_key = hash_token("sk-project-test-2")
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
|
||||
captured_descriptors = []
|
||||
|
||||
async def mock_should_rate_limit(descriptors, **kwargs):
|
||||
captured_descriptors.extend(descriptors)
|
||||
return {"overall_code": "OK", "statuses": []}
|
||||
|
||||
parallel_request_handler.should_rate_limit = mock_should_rate_limit
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key=_api_key,
|
||||
project_id="proj-abc123",
|
||||
project_metadata={
|
||||
"model_rpm_limit": {"gpt-4": 5},
|
||||
},
|
||||
)
|
||||
|
||||
# Request for gpt-3.5-turbo — project only limits gpt-4
|
||||
await parallel_request_handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=local_cache,
|
||||
data={"model": "gpt-3.5-turbo"},
|
||||
call_type="",
|
||||
)
|
||||
|
||||
descriptor_keys = [d["key"] for d in captured_descriptors]
|
||||
assert (
|
||||
"model_per_project" not in descriptor_keys
|
||||
), f"model_per_project should not be added for unrelated model, got: {descriptor_keys}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue