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:
shivam 2026-04-17 18:17:24 -07:00
parent 850fe595ac
commit 6fd49f1da1
No known key found for this signature in database
3 changed files with 171 additions and 1 deletions

View file

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

View file

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

View file

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