fix(rate_limiter): resolve model aliases and access groups for per-model limits

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
shivam 2026-09-04 20:42:55 +00:00
parent 338a37d8cd
commit 25733020ce
4 changed files with 276 additions and 69 deletions

View file

@ -1271,6 +1271,24 @@ def get_model_rate_limit_from_metadata(
return None
def resolve_rate_limited_model_name(
requested_model: str,
configured_models: Collection[str],
llm_router: Router | None,
) -> str | None:
"""The configured name that governs `requested_model`: itself, its model_group_alias target,
or a model access group serving it."""
if requested_model in configured_models:
return requested_model
if llm_router is None:
return None
candidates: Final = (
llm_router._get_model_from_alias(requested_model),
*llm_router.get_model_access_groups(model_name=requested_model),
)
return next((c for c in candidates if c is not None and c in configured_models), None)
def get_team_model_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> dict[str, int] | None:

View file

@ -37,6 +37,7 @@ from litellm.proxy.auth.auth_utils import (
get_estimated_output_tokens,
get_key_tag_rpm_limit,
get_model_rate_limit_from_metadata,
resolve_rate_limited_model_name,
)
from litellm.proxy.auth.budget_throttle import throttled_limit
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
@ -67,6 +68,7 @@ if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
from litellm.router import Router
from litellm.types.agents import AgentResponse
from litellm.types.caching import RedisPipelineIncrementOperation
@ -379,6 +381,12 @@ PROJECT_OTPM_DESCRIPTOR_KEY: Final = "model_per_project_otpm"
PARALLEL_REQUEST_SLOT_TTL_SECONDS: Final = 3600
def _proxy_llm_router() -> "Router | None":
from litellm.proxy.proxy_server import llm_router
return llm_router
CacheCounterValue: TypeAlias = int | float | str | bytes
CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None]
@ -2277,35 +2285,31 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
or get_model_rate_limit_from_metadata(user_api_key_dict, "organization_metadata", "model_tpm_limit")
is not None
):
_tpm_limit_for_team_model: Final = (
_tpm_limit_for_org_model: Final = (
get_model_rate_limit_from_metadata(user_api_key_dict, "organization_metadata", "model_tpm_limit") or {}
)
_rpm_limit_for_team_model: Final = (
_rpm_limit_for_org_model: Final = (
get_model_rate_limit_from_metadata(user_api_key_dict, "organization_metadata", "model_rpm_limit") or {}
)
should_check_rate_limit = False
if requested_model in _tpm_limit_for_team_model or requested_model in _rpm_limit_for_team_model:
should_check_rate_limit = True
if should_check_rate_limit:
model_specific_tpm_limit = None
model_specific_rpm_limit = None
if requested_model in _tpm_limit_for_team_model:
model_specific_tpm_limit = _tpm_limit_for_team_model[requested_model]
if requested_model in _rpm_limit_for_team_model:
model_specific_rpm_limit = _rpm_limit_for_team_model[requested_model]
descriptors.append(
RateLimitDescriptor(
key="model_per_organization",
value=f"{user_api_key_dict.org_id}:{requested_model}",
rate_limit={
"requests_per_unit": model_specific_rpm_limit,
"tokens_per_unit": model_specific_tpm_limit,
"window_size": self.window_size,
},
)
if requested_model is not None:
configured: Final = _tpm_limit_for_org_model.keys() | _rpm_limit_for_org_model.keys()
limited_model: Final = resolve_rate_limited_model_name(
requested_model=requested_model,
configured_models=configured,
llm_router=_proxy_llm_router(),
)
if limited_model is not None:
descriptors.append(
RateLimitDescriptor(
key="model_per_organization",
value=f"{user_api_key_dict.org_id}:{limited_model}",
rate_limit={
"requests_per_unit": _rpm_limit_for_org_model.get(limited_model),
"tokens_per_unit": _tpm_limit_for_org_model.get(limited_model),
"window_size": self.window_size,
},
)
)
return descriptors
@ -2340,22 +2344,23 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
_tpm_limit_for_key_model = _tpm_limit_for_key_model or {}
_rpm_limit_for_key_model = _rpm_limit_for_key_model or {}
# Check if model has any rate limits configured
should_check_rate_limit: Final = (
requested_model in _tpm_limit_for_key_model or requested_model in _rpm_limit_for_key_model
configured: Final = _tpm_limit_for_key_model.keys() | _rpm_limit_for_key_model.keys()
limited_model: Final = resolve_rate_limited_model_name(
requested_model=requested_model,
configured_models=configured,
llm_router=_proxy_llm_router(),
)
if not should_check_rate_limit:
if limited_model is None:
return
# Get model-specific limits
model_specific_tpm_limit: Final[int | None] = _tpm_limit_for_key_model.get(requested_model)
model_specific_rpm_limit: Final[int | None] = _rpm_limit_for_key_model.get(requested_model)
model_specific_tpm_limit: Final[int | None] = _tpm_limit_for_key_model.get(limited_model)
model_specific_rpm_limit: Final[int | None] = _rpm_limit_for_key_model.get(limited_model)
descriptors.append(
RateLimitDescriptor(
key="model_per_key",
value=f"{user_api_key_dict.api_key}:{requested_model}",
value=f"{user_api_key_dict.api_key}:{limited_model}",
rate_limit={
"requests_per_unit": model_specific_rpm_limit,
"tokens_per_unit": model_specific_tpm_limit,
@ -2900,24 +2905,25 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
_rpm_limit_for_team_model: Final = (
get_model_rate_limit_from_metadata(user_api_key_dict, "team_metadata", "model_rpm_limit") or {}
)
should_check_rate_limit: Final = (
requested_model in _tpm_limit_for_team_model or requested_model in _rpm_limit_for_team_model
)
if should_check_rate_limit and requested_model is not None:
model_specific_tpm_limit: Final = _tpm_limit_for_team_model.get(requested_model)
model_specific_rpm_limit: Final = _rpm_limit_for_team_model.get(requested_model)
descriptors.append(
RateLimitDescriptor(
key="model_per_team",
value=f"{user_api_key_dict.team_id}:{requested_model}",
rate_limit={
"requests_per_unit": model_specific_rpm_limit,
"tokens_per_unit": model_specific_tpm_limit,
"window_size": self.window_size,
},
)
if requested_model is not None:
configured: Final = _tpm_limit_for_team_model.keys() | _rpm_limit_for_team_model.keys()
limited_model: Final = resolve_rate_limited_model_name(
requested_model=requested_model,
configured_models=configured,
llm_router=_proxy_llm_router(),
)
if limited_model is not None:
descriptors.append(
RateLimitDescriptor(
key="model_per_team",
value=f"{user_api_key_dict.team_id}:{limited_model}",
rate_limit={
"requests_per_unit": _rpm_limit_for_team_model.get(limited_model),
"tokens_per_unit": _tpm_limit_for_team_model.get(limited_model),
"window_size": self.window_size,
},
)
)
def _add_project_model_rate_limit_descriptor_from_metadata(
self,
@ -2936,24 +2942,25 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
_rpm_limit_for_project_model: Final = (
get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_rpm_limit") or {}
)
should_check_rate_limit: Final = (
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: Final = _tpm_limit_for_project_model.get(requested_model)
model_specific_rpm_limit: Final = _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,
},
)
if requested_model is not None:
configured: Final = _tpm_limit_for_project_model.keys() | _rpm_limit_for_project_model.keys()
limited_model: Final = resolve_rate_limited_model_name(
requested_model=requested_model,
configured_models=configured,
llm_router=_proxy_llm_router(),
)
if limited_model is not None:
descriptors.append(
RateLimitDescriptor(
key="model_per_project",
value=f"{user_api_key_dict.project_id}:{limited_model}",
rate_limit={
"requests_per_unit": _rpm_limit_for_project_model.get(limited_model),
"tokens_per_unit": _tpm_limit_for_project_model.get(limited_model),
"window_size": self.window_size,
},
)
)
def add_project_io_token_rate_limit_descriptors_from_metadata(
self,
@ -2979,13 +2986,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
or {} # mutable-ok: metadata helper returns an optional mapping
)
model_itpm_limit: Final = itpm_limit_for_project_model.get(requested_model)
model_otpm_limit: Final = otpm_limit_for_project_model.get(requested_model)
configured: Final = itpm_limit_for_project_model.keys() | otpm_limit_for_project_model.keys()
limited_model: Final = resolve_rate_limited_model_name(
requested_model=requested_model,
configured_models=configured,
llm_router=_proxy_llm_router(),
)
if limited_model is None:
return
model_itpm_limit: Final = itpm_limit_for_project_model.get(limited_model)
model_otpm_limit: Final = otpm_limit_for_project_model.get(limited_model)
if model_itpm_limit is None and model_otpm_limit is None:
return
descriptor_value: Final = f"{user_api_key_dict.project_id}:{requested_model}"
descriptor_value: Final = f"{user_api_key_dict.project_id}:{limited_model}"
if model_itpm_limit is not None:
descriptors.append(
RateLimitDescriptor(

View file

@ -26,9 +26,21 @@ from litellm.proxy.auth.auth_utils import (
get_project_model_tpm_limit,
get_request_route_template,
is_request_body_safe,
resolve_rate_limited_model_name,
)
@pytest.mark.parametrize(
("requested_model", "configured_models", "expected"),
[
("gpt-4", {"gpt-4"}, "gpt-4"),
("gpt4", {"gpt-4"}, None),
],
)
def test_resolve_rate_limited_model_name_without_router(requested_model, configured_models, expected):
assert resolve_rate_limited_model_name(requested_model, configured_models, None) == expected
class TestCustomAuthCommonChecksWarning:
"""custom_auth_common_checks_warning only warns when custom auth is configured
and the common-checks opt-in is off, since that is the only state where

View file

@ -51,6 +51,27 @@ class TimeController:
self._current += timedelta(seconds=seconds)
@pytest.fixture
def model_rate_limit_router() -> Router:
return Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "openai/gpt-4"},
"model_info": {"access_groups": ["gpt4-family"]},
},
{
"model_name": "gpt-4-turbo",
"litellm_params": {"model": "openai/gpt-4-turbo"},
"model_info": {"access_groups": ["gpt4-family"]},
},
],
model_group_alias={"gpt4": "gpt-4"},
set_verbose=False,
num_retries=0,
)
@pytest.fixture
def time_controller(monkeypatch):
controller = TimeController()
@ -65,6 +86,147 @@ def _isolated_request_stash():
_request_stash.reset(token)
@pytest.mark.parametrize(
("metadata", "requested_model", "expected_model", "expected_tokens"),
[
({"model_tpm_limit": {"gpt-4": 100}}, "gpt4", "gpt-4", 100),
({"model_tpm_limit": {"gpt4-family": 100}}, "gpt-4-turbo", "gpt4-family", 100),
(
{"model_tpm_limit": {"gpt-4": 100, "gpt4": 50}},
"gpt4",
"gpt4",
50,
),
({"model_tpm_limit": {"other-model": 100}}, "gpt4", None, None),
],
)
def test_model_per_key_rate_limit_resolves_alias_and_access_group(
monkeypatch,
model_rate_limit_router,
metadata,
requested_model,
expected_model,
expected_tokens,
):
monkeypatch.setattr(
"litellm.proxy.hooks.parallel_request_limiter_v3._proxy_llm_router",
lambda: model_rate_limit_router,
)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = handler._create_rate_limit_descriptors(
user_api_key_dict=UserAPIKeyAuth(api_key="sk-model-key", metadata=metadata),
data={"model": requested_model},
rpm_limit_type=None,
tpm_limit_type=None,
model_has_failures=False,
)
model_descriptor = next(
(descriptor for descriptor in descriptors if descriptor["key"] == "model_per_key"),
None,
)
if expected_model is None:
assert model_descriptor is None
return
assert model_descriptor is not None
assert model_descriptor["value"] == f"{hash_token('sk-model-key')}:{expected_model}"
assert model_descriptor["rate_limit"]["tokens_per_unit"] == expected_tokens
def test_model_per_team_rate_limit_resolves_alias(monkeypatch, model_rate_limit_router):
monkeypatch.setattr(
"litellm.proxy.hooks.parallel_request_limiter_v3._proxy_llm_router",
lambda: model_rate_limit_router,
)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = []
handler._add_team_model_rate_limit_descriptor_from_metadata(
user_api_key_dict=UserAPIKeyAuth(
team_id="team-model-key",
team_metadata={"model_tpm_limit": {"gpt-4": 100}},
),
requested_model="gpt4",
descriptors=descriptors,
)
assert descriptors == [
{
"key": "model_per_team",
"value": "team-model-key:gpt-4",
"rate_limit": {
"requests_per_unit": None,
"tokens_per_unit": 100,
"window_size": handler.window_size,
},
}
]
@pytest.mark.asyncio
async def test_model_per_key_rate_limit_alias_shares_tpm_counter(
monkeypatch, model_rate_limit_router
):
monkeypatch.setattr(
"litellm.proxy.hooks.parallel_request_limiter_v3._proxy_llm_router",
lambda: model_rate_limit_router,
)
api_key = hash_token("sk-model-counter")
user_api_key_dict = UserAPIKeyAuth(
api_key=api_key,
metadata={"model_tpm_limit": {"gpt-4": 10}},
)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
first_request = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 5,
}
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data=first_request,
call_type="",
)
await handler.async_log_success_event(
kwargs={
"standard_logging_object": {
"metadata": {"user_api_key_hash": api_key},
},
"model": "gpt-4",
},
response_obj=ModelResponse(
id="model-counter",
object="chat.completion",
created=int(datetime.now().timestamp()),
model="gpt-4",
usage=Usage(prompt_tokens=5, completion_tokens=5, total_tokens=10),
choices=[],
),
start_time=datetime.now(),
end_time=datetime.now(),
)
counter_key = handler.create_rate_limit_keys(
"model_per_key",
f"{api_key}:gpt-4",
"tokens",
)
assert await local_cache.async_get_cache(key=counter_key) == 10
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"model": "gpt4"},
call_type="",
)
assert exc_info.value.status_code == 429
@pytest.mark.parametrize(
"throttle_pct, expected_rpm, expected_tpm",
[