fix(proxy): let admins make per-model key/team rate limits a hard 429 instead of falling back (#45459)

* fix(proxy): let admins make per-model key/team rate limits a hard 429 instead of falling back

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): sync per-model rate limit API types

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): type per-model fallback setting access

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep the per-model hard cap inside the fallback loop and test it through the real limiter

* fix(proxy): list the per-model hard cap setting on the settings page and cover every per-model descriptor

* fix(proxy): read disable_fallbacks_on_per_model_rate_limits as a boolean value and cover the team-wide per-model cap

* test(proxy): cover per-model rate limit 429s skipping fallbacks on a live proxy

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
Shivam Rawat 2026-10-09 23:10:23 -07:00 • committed by GitHub
parent ed9c31cfa8
commit 15154a40a8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 1242 additions and 13 deletions

View file

@ -3130,6 +3130,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="If True, router fallbacks configured in router_settings are only attempted when the calling key (and its team and project) is allowed to call the fallback model; unauthorized fallback targets are skipped and the primary model's error is returned. Default is False.",
)
disable_fallbacks_on_per_model_rate_limits: bool | None = Field(
None,
description=(
"If true, a request rejected by a key/team/org/project per-model rate limit "
"(model_rpm_limit / model_tpm_limit) returns 429 instead of retrying on the "
"configured fallbacks"
),
)
scheduled_job_stagger: ScheduledJobStaggerSettings | None = Field(
None,
description=(

View file

@ -16,6 +16,7 @@ from typing import (
Protocol,
TypeAlias,
TypeVar,
cast, # noqa: TID251 # _pre_call_with_fallbacks receives general_settings as a legacy bare dict
overload,
runtime_checkable,
)
@ -2302,7 +2303,21 @@ class ProxyBaseLLMRequestProcessing:
route_type: str,
llm_router: Router | None,
) -> tuple[dict, LiteLLMLoggingObj]:
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.common_utils.proxy_rate_limit_error import (
PER_MODEL_RATE_LIMIT_DESCRIPTOR_KEYS,
ProxyRateLimitError,
per_model_rate_limits_disable_fallbacks,
)
general_settings_view: Final = cast( # cast-ok: this method keeps its legacy bare-dict settings parameter
Mapping[str, object], general_settings
)
def is_hard_per_model_limit(exc: ProxyRateLimitError) -> bool:
return (
exc.descriptor_key in PER_MODEL_RATE_LIMIT_DESCRIPTOR_KEYS
and per_model_rate_limits_disable_fallbacks(general_settings_view)
)
configured_fallbacks: Final = (
self._configured_fallbacks(llm_router=llm_router, user_api_key_dict=user_api_key_dict)
@ -2336,6 +2351,7 @@ class ProxyBaseLLMRequestProcessing:
or not configured_fallbacks
or rate_limited_data.get("disable_fallbacks")
or not isinstance(original_model, str)
or is_hard_per_model_limit(original_exc)
):
raise
@ -2375,7 +2391,9 @@ class ProxyBaseLLMRequestProcessing:
llm_router=llm_router,
rate_limited_model=original_model,
)
except ProxyRateLimitError:
except ProxyRateLimitError as fallback_exc:
if is_hard_per_model_limit(fallback_exc):
raise
continue
except BaseException:
self.data = rate_limited_data

View file

@ -40,9 +40,37 @@ from collections.abc import Mapping
from typing import Any, Final
from fastapi import HTTPException
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import RateLimitError, RateLimitErrorCategory, RateLimitType
PER_MODEL_RATE_LIMIT_DESCRIPTOR_KEYS: Final = frozenset(
{
"model_per_key",
"model_per_team",
"model_per_organization",
"model_per_project",
"model_per_project_itpm",
"model_per_project_otpm",
}
)
DISABLE_FALLBACKS_ON_PER_MODEL_RATE_LIMITS_SETTING: Final = "disable_fallbacks_on_per_model_rate_limits"
_DISABLE_FALLBACKS_ON_PER_MODEL_RATE_LIMITS_FLAG: Final[TypeAdapter[bool | None]] = TypeAdapter(bool | None)
def per_model_rate_limits_disable_fallbacks(general_settings: Mapping[str, object]) -> bool:
raw_value: Final = general_settings.get(DISABLE_FALLBACKS_ON_PER_MODEL_RATE_LIMITS_SETTING)
try:
return _DISABLE_FALLBACKS_ON_PER_MODEL_RATE_LIMITS_FLAG.validate_python(raw_value) is True
except ValidationError:
verbose_proxy_logger.warning(
"general_settings.%s=%r is not a boolean, treating it as disabled",
DISABLE_FALLBACKS_ON_PER_MODEL_RATE_LIMITS_SETTING,
raw_value,
)
return False
def map_v3_rate_limit_type(
v3_value: str | None,
@ -149,6 +177,8 @@ class ProxyRateLimitError(HTTPException, RateLimitError):
rate_limit_type: str | RateLimitType | None = None,
model: str | None = None,
llm_provider: str | None = "litellm_proxy",
*,
descriptor_key: str | None = None,
):
# Normalize None → safe defaults so callers (and the resolver helper
# in `rate_limiter_utils`) can pass `None` without producing an
@ -191,3 +221,4 @@ class ProxyRateLimitError(HTTPException, RateLimitError):
self.headers = stringified_headers
self.detail = detail
self.status_code = 429
self.descriptor_key = descriptor_key

View file

@ -105,6 +105,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
additional_details=f"{CommonProxyErrors.max_parallel_request_limit_reached.value}. Hit limit for {rate_limit_type}. Current limits: max_parallel_requests: {max_parallel_requests}, tpm_limit: {tpm_limit}, rpm_limit: {rpm_limit}",
rate_limit_type=triggered_type,
requested_model=data.get("model") if data else None,
descriptor_key=rate_limit_type,
)
new_val = {
"current_requests": 1,
@ -143,6 +144,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
rate_limit_type=triggered_type,
model=resolved_model,
llm_provider=llm_provider,
descriptor_key=rate_limit_type,
)
await self.internal_usage_cache.async_batch_set_cache(
@ -170,6 +172,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
additional_details: str | None = None,
rate_limit_type: RateLimitType | None = None,
requested_model: str | None = None,
descriptor_key: str | None = None,
) -> NoReturn:
"""
Raise a 429 with a retry-after header for litellm-proxy parallel-request limits.
@ -207,6 +210,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
rate_limit_type=rate_limit_type or RateLimitType.CONCURRENT_REQUESTS,
model=resolved_model,
llm_provider=llm_provider,
descriptor_key=descriptor_key,
)
@with_service_target("rate_limits")

View file

@ -3594,6 +3594,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
rate_limit_type=map_v3_rate_limit_type(status["rate_limit_type"]),
model=resolved_model,
llm_provider=llm_provider,
descriptor_key=descriptor_key,
)
@staticmethod

View file

@ -18399,6 +18399,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
"pass_through_endpoints": "PydanticModel",
"store_model_in_db": "Boolean",
"store_prompts_in_spend_logs": "Boolean",
"disable_fallbacks_on_per_model_rate_limits": "Boolean",
"maximum_spend_logs_retention_period": "String",
"maximum_health_check_retention_period": "String",
"maximum_daily_tag_spend_retention_period": "String",

View file

@ -0,0 +1,647 @@
from __future__ import annotations
import json
import re
import signal
import uuid
from collections import Counter
from collections.abc import Iterator, Mapping, Sequence
from concurrent.futures import ThreadPoolExecutor
from contextlib import ExitStack
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
import anthropic
import httpx
import openai
import psutil
import pytest
import yaml
from integration._support.client import (
JSON_OBJECT,
Gateway,
Scenario,
eventually,
gateway_from_environment,
object_value,
string_value,
)
from integration._support.openai_wire import chat_reply, responses_reply
from integration._support.process import graceful_stop_seconds, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.routing.test_payment_required_mapping import OPENAI_MODEL, is_model_info_probe, model_list_reply
from pydantic import JsonValue, TypeAdapter
pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
SETTING: Final = "disable_fallbacks_on_per_model_rate_limits"
FRESH_CONNECTION: Final = MappingProxyType({"Connection": "close"})
MARKER: Final = re.compile(rb"pmrl-[0-9a-f]{32}")
WORKER_STARTED: Final = re.compile(r"Started server process \[(\d+)\]")
STARTUP_COMPLETE: Final = "Application startup complete."
CONFIG_LIST: Final = TypeAdapter(list[dict[str, JsonValue]])
WORKER_PIDS: Final = TypeAdapter(tuple[int, ...])
BURST: Final = 8
TOKEN_LIMIT: Final = 100
INPUT_TOKEN_LIMIT: Final = 120
MAX_TOKENS: Final = 60
METERED_USAGE: Final = MappingProxyType({"prompt_tokens": 900, "completion_tokens": 100, "total_tokens": 1000})
MALFORMED_VALUES: Final[tuple[tuple[str, JsonValue], ...]] = (
("word", "sometimes"),
("empty", ""),
("five_kb", "x" * 5120),
("list", [True]),
("dict", {"enabled": True}),
("int", 2),
)
@dataclass(frozen=True, slots=True)
class Groups:
primary: str
fallback: str
last: str
metered: str
gated_hard: str
gated_soft: str
capped: str
def names(self) -> tuple[str, ...]:
return (self.primary, self.fallback, self.last, self.metered, self.gated_hard, self.gated_soft, self.capped)
@dataclass(frozen=True, slots=True)
class Rig:
gateway: Gateway
log: Path
groups: Groups
wires: Mapping[str, Wire]
def traffic(self) -> dict[str, tuple[str, ...]]:
return {group: markers_on(wire) for group, wire in self.wires.items()}
def expected(self, served: Mapping[str, Sequence[str]]) -> dict[str, tuple[str, ...]]:
return {group: tuple(sorted(served.get(group, ()))) for group in self.wires}
def new_marker() -> str:
return f"pmrl-{uuid.uuid4().hex}"
def new_groups() -> Groups:
suffix: Final = uuid.uuid4().hex[:10]
return Groups(
primary=f"per-model-primary-{suffix}",
fallback=f"per-model-fallback-{suffix}",
last=f"per-model-last-{suffix}",
metered=f"per-model-metered-{suffix}",
gated_hard=f"per-model-gated-hard-{suffix}",
gated_soft=f"per-model-gated-soft-{suffix}",
capped=f"per-model-capped-{suffix}",
)
def marker_in(request: Request) -> str:
found: Final = MARKER.search(request.body)
assert found is not None, (request.method, request.target, request.body[:300])
return found.group(0).decode()
def markers_on(wire: Wire) -> tuple[str, ...]:
return tuple(sorted(marker_in(request) for request in wire.drain() if not is_model_info_probe(request)))
def streamed(request: Request) -> bool:
return JSON_OBJECT.validate_json(request.body).get("stream") is True
def healthy_reply(request: Request) -> Reply:
if is_model_info_probe(request):
return model_list_reply()
marker: Final = marker_in(request)
if request.target.endswith("/responses"):
return responses_reply(f"resp_{marker}", OPENAI_MODEL, f"served {marker}", stream=streamed(request))
return chat_reply(f"chatcmpl-{marker}", OPENAI_MODEL, f"served {marker}", stream=streamed(request))
def metered_reply(request: Request) -> Reply:
if is_model_info_probe(request):
return model_list_reply()
marker: Final = marker_in(request)
body: Final = {
"id": f"chatcmpl-{marker}",
"object": "chat.completion",
"created": 1,
"model": OPENAI_MODEL,
"choices": [
{"index": 0, "message": {"role": "assistant", "content": f"served {marker}"}, "finish_reason": "stop"}
],
"usage": dict(METERED_USAGE),
}
return Reply(body=json.dumps(body).encode())
def model_entry(group: str, wire: Wire, **extra: JsonValue) -> dict[str, JsonValue]:
return {
"model_name": group,
"litellm_params": {
"model": OPENAI_MODEL,
"api_base": wire.url + "/v1",
"api_key": "synthetic-openai-key",
**extra,
},
}
def proxy_config(
directory: Path,
*,
model_list: Sequence[Mapping[str, JsonValue]],
fallbacks: Sequence[Mapping[str, JsonValue]],
general_settings: Mapping[str, JsonValue],
callbacks: Sequence[str] = (),
) -> Path:
base: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
litellm_settings: Final = object_value(base["litellm_settings"])
config: Final = {
**base,
"model_list": [dict(entry) for entry in model_list],
"general_settings": {**object_value(base["general_settings"]), **general_settings},
"litellm_settings": {**litellm_settings, **({"callbacks": list(callbacks)} if callbacks else {})},
"router_settings": {
**object_value(base["router_settings"]),
"num_retries": 0,
"fallbacks": [dict(entry) for entry in fallbacks],
},
}
path: Final = directory / f"per-model-limits-{uuid.uuid4().hex[:8]}.yaml"
path.write_text(yaml.safe_dump(config))
return path
def enforcing_config(directory: Path, groups: Groups, wires: Mapping[str, Wire]) -> Path:
return proxy_config(
directory,
model_list=(
model_entry(groups.primary, wires[groups.primary]),
model_entry(groups.fallback, wires[groups.fallback]),
model_entry(groups.last, wires[groups.last]),
model_entry(groups.metered, wires[groups.metered]),
model_entry(groups.gated_hard, wires[groups.gated_hard], rpm=1),
model_entry(groups.gated_soft, wires[groups.gated_soft], rpm=1),
model_entry(groups.capped, wires[groups.capped]),
),
fallbacks=(
{groups.primary: [groups.fallback]},
{groups.metered: [groups.fallback]},
{groups.gated_hard: [groups.capped, groups.last]},
{groups.gated_soft: [groups.fallback]},
),
general_settings={SETTING: True},
callbacks=("dynamic_rate_limiter_v3",),
)
@pytest.fixture(scope="module")
def enforcing(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
directory: Final = tmp_path_factory.mktemp("per-model-limits")
groups: Final = new_groups()
with ExitStack() as stack:
rig_gateway: Final = stack.enter_context(gateway_from_environment())
wires: Final = MappingProxyType(
{
group: stack.enter_context(wire_server(metered_reply if group == groups.metered else healthy_reply))
for group in groups.names()
}
)
owned: Final = stack.enter_context(
owned_proxy_process(
rig_gateway, directory, {}, config=enforcing_config(directory, groups, wires), workers=2
)
)
yield Rig(owned.gateway, owned.log, groups, wires)
def chat(gateway: Gateway, key: str, group: str, content: str, **extra: JsonValue) -> httpx.Response:
return gateway.request(
"POST",
"/v1/chat/completions",
{"model": group, "messages": [{"role": "user", "content": content}], **extra},
key=key,
headers=FRESH_CONNECTION,
)
def responses_call(gateway: Gateway, key: str, group: str, marker: str, *, stream: bool) -> httpx.Response:
return gateway.request(
"POST", "/v1/responses", {"model": group, "input": marker, "stream": stream}, key=key, headers=FRESH_CONNECTION
)
def openai_sdk(gateway: Gateway, key: str) -> openai.OpenAI:
return openai.OpenAI(
base_url=str(gateway.client.base_url).rstrip("/") + "/v1",
api_key=key,
max_retries=0,
default_headers=dict(FRESH_CONNECTION),
)
def openai_async_sdk(gateway: Gateway, key: str) -> openai.AsyncOpenAI:
return openai.AsyncOpenAI(
base_url=str(gateway.client.base_url).rstrip("/") + "/v1",
api_key=key,
max_retries=0,
default_headers=dict(FRESH_CONNECTION),
)
def anthropic_sdk(gateway: Gateway, key: str) -> anthropic.Anthropic:
return anthropic.Anthropic(
base_url=str(gateway.client.base_url).rstrip("/"),
api_key=key,
max_retries=0,
default_headers=dict(FRESH_CONNECTION),
)
def anthropic_async_sdk(gateway: Gateway, key: str) -> anthropic.AsyncAnthropic:
return anthropic.AsyncAnthropic(
base_url=str(gateway.client.base_url).rstrip("/"),
api_key=key,
max_retries=0,
default_headers=dict(FRESH_CONNECTION),
)
def error_message(body: str) -> str:
return string_value(object_value(JSON_OBJECT.validate_json(body)["error"])["message"])
def assert_limit_message(body: str, descriptor: str, group: str, kind: str, limit: int) -> None:
message: Final = error_message(body)
assert f"Rate limit exceeded for {descriptor}: " in message, message
assert f":{group}. Limit type: {kind}. Current limit: {limit}, Remaining: " in message, message
def assert_refused(
response: httpx.Response, descriptor: str, group: str, *, kind: str = "requests", limit: int = 1
) -> None:
assert response.status_code == 429, (
response.status_code,
response.headers.get("x-litellm-model-group"),
response.text,
)
assert_limit_message(response.text, descriptor, group, kind, limit)
def assert_served(response: httpx.Response, group: str, identity: str) -> None:
assert response.status_code == 200, response.text
assert response.headers["x-litellm-model-group"] == group, (
response.headers.get("x-litellm-model-group"),
response.text,
)
assert identity in response.text, response.text
def burst(gateway: Gateway, calls: Sequence[tuple[str, str]], group: str) -> tuple[httpx.Response, ...]:
def send(call: tuple[str, str]) -> httpx.Response:
return chat(gateway, call[0], group, call[1])
with ThreadPoolExecutor(max_workers=len(calls)) as pool:
return tuple(pool.map(send, calls))
def served_once_in_a_burst(rig: Rig, markers: Sequence[str], responses: Sequence[httpx.Response]) -> str:
assert Counter(response.status_code for response in responses) == Counter({200: 1, 429: BURST - 1}), tuple(
(response.status_code, response.headers.get("x-litellm-model-group"), response.text[:200])
for response in responses
)
refusals: Final = tuple(response for response in responses if response.status_code == 429)
for response in refusals:
assert_refused(response, "model_per_key", rig.groups.primary)
(served,) = (marker for marker, response in zip(markers, responses) if response.status_code == 200)
return served
def worker_startups(log: Path) -> tuple[tuple[int, ...], int]:
text: Final = log.read_text(errors="replace")
return WORKER_PIDS.validate_python(WORKER_STARTED.findall(text)), text.count(STARTUP_COMPLETE)
def listed_setting(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]:
response: Final = gateway.request("GET", "/config/list", params={"config_type": "general_settings"})
assert response.status_code == 200, response.text
return tuple(entry for entry in CONFIG_LIST.validate_json(response.content) if entry["field_name"] == SETTING)
def update_setting(gateway: Gateway, value: JsonValue) -> httpx.Response:
return gateway.request(
"POST",
"/config/field/update",
{"field_name": SETTING, "field_value": value, "config_type": "general_settings"},
)
def scope_key_fields(scenario: Scenario, scope: str, group: str) -> dict[str, JsonValue]:
limit: Final[dict[str, JsonValue]] = {group: 1}
if scope == "team":
return {"team_id": scenario.team(model_rpm_limit=limit)}
if scope == "organization":
return {"team_id": scenario.team(organization_id=scenario.organization(model_rpm_limit=limit))}
team: Final = scenario.team()
return {"team_id": team, "project_id": scenario.project(team, model_rpm_limit=limit)}
def test_chat_openai_sdk_over_key_model_rpm_gets_429_instead_of_the_fallback(enforcing: Rig) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
client: Final = openai_sdk(enforcing.gateway, scenario.key(model_rpm_limit={primary: 1}))
served: Final = client.chat.completions.with_raw_response.create(
model=primary, messages=[{"role": "user", "content": first}]
)
assert served.headers["x-litellm-model-group"] == primary
assert served.parse().id == f"chatcmpl-{first}"
with pytest.raises(openai.RateLimitError) as refused:
client.chat.completions.create(model=primary, messages=[{"role": "user", "content": second}])
assert_limit_message(refused.value.response.text, "model_per_key", primary, "requests", 1)
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
async def test_chat_stream_openai_async_sdk_over_key_model_rpm_gets_429_instead_of_the_fallback(
enforcing: Rig,
) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
async with openai_async_sdk(enforcing.gateway, scenario.key(model_rpm_limit={primary: 1})) as client:
stream: Final = await client.chat.completions.create(
model=primary, messages=[{"role": "user", "content": first}], stream=True
)
identities: Final = {chunk.id async for chunk in stream}
assert identities == {f"chatcmpl-{first}"}
with pytest.raises(openai.RateLimitError) as refused:
await client.chat.completions.create(
model=primary, messages=[{"role": "user", "content": second}], stream=True
)
assert_limit_message(refused.value.response.text, "model_per_key", primary, "requests", 1)
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
def test_messages_anthropic_sdk_over_key_model_rpm_gets_429_instead_of_the_fallback(enforcing: Rig) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
client: Final = anthropic_sdk(enforcing.gateway, scenario.key(model_rpm_limit={primary: 1}))
served: Final = client.messages.with_raw_response.create(
model=primary, max_tokens=32, messages=[{"role": "user", "content": first}]
)
assert served.headers["x-litellm-model-group"] == primary
block: Final = served.parse().content[0]
assert isinstance(block, anthropic.types.TextBlock), block
assert block.text == f"served {first}"
with pytest.raises(anthropic.RateLimitError) as refused:
client.messages.create(model=primary, max_tokens=32, messages=[{"role": "user", "content": second}])
assert_limit_message(refused.value.response.text, "model_per_key", primary, "requests", 1)
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
async def test_messages_stream_anthropic_async_sdk_over_key_model_rpm_gets_429_instead_of_the_fallback(
enforcing: Rig,
) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
async with anthropic_async_sdk(enforcing.gateway, scenario.key(model_rpm_limit={primary: 1})) as client:
async with client.messages.stream(
model=primary, max_tokens=32, messages=[{"role": "user", "content": first}]
) as served:
text: Final = await served.get_final_text()
assert text == f"served {first}"
with pytest.raises(anthropic.RateLimitError) as refused:
async with client.messages.stream(
model=primary, max_tokens=32, messages=[{"role": "user", "content": second}]
) as rejected:
await rejected.get_final_text()
assert_limit_message(refused.value.response.text, "model_per_key", primary, "requests", 1)
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
@pytest.mark.parametrize("stream", (False, True), ids=("json", "stream"))
def test_responses_httpx_over_key_model_rpm_gets_429_instead_of_the_fallback(enforcing: Rig, stream: bool) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
key: Final = scenario.key(model_rpm_limit={primary: 1})
assert_served(responses_call(enforcing.gateway, key, primary, first, stream=stream), primary, f"resp_{first}")
assert_refused(responses_call(enforcing.gateway, key, primary, second, stream=stream), "model_per_key", primary)
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
@pytest.mark.parametrize(
("scope", "descriptor"),
(("team", "model_per_team"), ("organization", "model_per_organization"), ("project", "model_per_project")),
ids=("team", "organization", "project"),
)
def test_chat_over_a_scope_model_rpm_gets_429_instead_of_the_fallback(
enforcing: Rig, scope: str, descriptor: str
) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
fields: Final = scope_key_fields(scenario, scope, primary)
first_key, second_key = scenario.key(**fields), scenario.key(**fields)
assert_served(chat(enforcing.gateway, first_key, primary, first), primary, f"chatcmpl-{first}")
assert_refused(chat(enforcing.gateway, second_key, primary, second), descriptor, primary)
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
def test_chat_over_project_model_itpm_gets_429_instead_of_the_fallback(enforcing: Rig) -> None:
enforcing.traffic()
metered: Final = enforcing.groups.metered
first, second = new_marker(), new_marker()
filler: Final = " word" * 50
with enforcing.gateway.scenario() as scenario:
team: Final = scenario.team()
key: Final = scenario.key(
team_id=team, project_id=scenario.project(team, model_itpm_limit={metered: INPUT_TOKEN_LIMIT})
)
assert_served(chat(enforcing.gateway, key, metered, first + filler), metered, f"chatcmpl-{first}")
assert_refused(
chat(enforcing.gateway, key, metered, second + filler),
"model_per_project_itpm",
metered,
kind="tokens",
limit=INPUT_TOKEN_LIMIT,
)
assert enforcing.traffic() == enforcing.expected({metered: (first,)})
def test_chat_over_key_model_tpm_gets_429_instead_of_the_fallback(enforcing: Rig) -> None:
enforcing.traffic()
metered: Final = enforcing.groups.metered
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
key: Final = scenario.key(model_tpm_limit={metered: TOKEN_LIMIT})
assert_served(chat(enforcing.gateway, key, metered, first, max_tokens=MAX_TOKENS), metered, f"chatcmpl-{first}")
assert_refused(
chat(enforcing.gateway, key, metered, second, max_tokens=MAX_TOKENS),
"model_per_key",
metered,
kind="tokens",
limit=TOKEN_LIMIT,
)
assert enforcing.traffic() == enforcing.expected({metered: (first,)})
def test_key_router_settings_fallbacks_are_skipped_on_a_per_model_429(enforcing: Rig) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
key: Final = scenario.key(
model_rpm_limit={primary: 1}, router_settings={"fallbacks": [{primary: [enforcing.groups.last]}]}
)
assert_served(chat(enforcing.gateway, key, primary, first), primary, f"chatcmpl-{first}")
assert_refused(chat(enforcing.gateway, key, primary, second), "model_per_key", primary)
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
def test_capacity_429_falls_back_but_a_capped_fallback_answers_429_instead_of_the_next_one(enforcing: Rig) -> None:
enforcing.traffic()
groups: Final = enforcing.groups
warm, spend, refused = new_marker(), new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
key: Final = scenario.key(model_rpm_limit={groups.capped: 1})
assert_served(chat(enforcing.gateway, key, groups.gated_hard, warm), groups.gated_hard, f"chatcmpl-{warm}")
assert_served(chat(enforcing.gateway, key, groups.capped, spend), groups.capped, f"chatcmpl-{spend}")
assert_refused(chat(enforcing.gateway, key, groups.gated_hard, refused), "model_per_key", groups.capped)
assert enforcing.traffic() == enforcing.expected({groups.gated_hard: (warm,), groups.capped: (spend,)})
def test_model_capacity_429_without_a_per_model_limit_still_falls_back(enforcing: Rig) -> None:
enforcing.traffic()
groups: Final = enforcing.groups
warm, moved = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
key: Final = scenario.key()
assert_served(chat(enforcing.gateway, key, groups.gated_soft, warm), groups.gated_soft, f"chatcmpl-{warm}")
assert_served(chat(enforcing.gateway, key, groups.gated_soft, moved), groups.fallback, f"chatcmpl-{moved}")
assert enforcing.traffic() == enforcing.expected({groups.gated_soft: (warm,), groups.fallback: (moved,)})
def test_key_rpm_limit_without_a_model_scope_answers_429_on_both_the_primary_and_the_fallback(enforcing: Rig) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
key: Final = scenario.key(rpm_limit=1)
assert_served(chat(enforcing.gateway, key, primary, first), primary, f"chatcmpl-{first}")
refused: Final = chat(enforcing.gateway, key, primary, second)
assert refused.status_code == 429, refused.text
assert error_message(refused.text).startswith("Rate limit exceeded for api_key: "), refused.text
assert ". Limit type: requests. Current limit: 1, Remaining: 0." in refused.text, refused.text
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
def test_response_cache_hits_count_toward_key_model_rpm_and_the_next_request_gets_429(enforcing: Rig) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
marker: Final = new_marker()
with enforcing.gateway.scenario() as scenario:
key: Final = scenario.key(model_rpm_limit={primary: 2})
assert_served(chat(enforcing.gateway, key, primary, marker), primary, f"chatcmpl-{marker}")
assert_served(chat(enforcing.gateway, key, primary, marker), primary, f"chatcmpl-{marker}")
assert_refused(chat(enforcing.gateway, key, primary, marker), "model_per_key", primary, limit=2)
assert enforcing.traffic() == enforcing.expected({primary: (marker,)})
def test_concurrent_burst_over_key_model_rpm_serves_one_and_refuses_the_rest_without_fallback(enforcing: Rig) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
markers: Final = tuple(new_marker() for _ in range(BURST))
bystander: Final = new_marker()
with enforcing.gateway.scenario() as scenario:
limited: Final = scenario.key(model_rpm_limit={primary: 1})
unlimited: Final = scenario.key()
calls: Final = (*((limited, marker) for marker in markers), (unlimited, bystander))
responses: Final = burst(enforcing.gateway, calls, primary)
served: Final = served_once_in_a_burst(enforcing, markers, responses[:BURST])
assert_served(responses[BURST], primary, f"chatcmpl-{bystander}")
assert enforcing.traffic() == enforcing.expected({primary: (served, bystander)})
def test_yaml_owned_setting_refuses_a_database_write_and_keeps_enforcing(enforcing: Rig) -> None:
enforcing.traffic()
primary: Final = enforcing.groups.primary
refused_write: Final = update_setting(enforcing.gateway, False)
assert refused_write.status_code == 400, refused_write.text
detail: Final = object_value(JSON_OBJECT.validate_json(refused_write.content)["detail"])
assert detail["error"] == (
f"general_settings key '{SETTING}' is set in the config file and cannot be changed here."
), detail
assert detail["keys"] == [SETTING], detail
listed: Final = listed_setting(enforcing.gateway)
assert [(entry["field_type"], entry["field_value"], entry["editable"]) for entry in listed] == [
("Boolean", True, False)
], listed
first, second = new_marker(), new_marker()
with enforcing.gateway.scenario() as scenario:
key: Final = scenario.key(model_rpm_limit={primary: 1})
assert_served(chat(enforcing.gateway, key, primary, first), primary, f"chatcmpl-{first}")
assert_refused(chat(enforcing.gateway, key, primary, second), "model_per_key", primary)
assert enforcing.traffic() == enforcing.expected({primary: (first,)})
def test_worker_sigkill_and_replacement_keeps_refusing_per_model_overflow_without_fallback(enforcing: Rig) -> None:
enforcing.traffic()
pids, ready = eventually(
lambda: worker_startups(enforcing.log), lambda found: len(found[0]) >= 2 and found[1] >= 2, seconds=60
)
victim: Final = next(pid for pid in pids if psutil.pid_exists(pid))
psutil.Process(victim).send_signal(signal.SIGKILL)
eventually(
lambda: worker_startups(enforcing.log),
lambda found: len(found[0]) > len(pids) and found[1] > ready,
seconds=graceful_stop_seconds(),
)
primary: Final = enforcing.groups.primary
markers: Final = tuple(new_marker() for _ in range(BURST))
with enforcing.gateway.scenario() as scenario:
limited: Final = scenario.key(model_rpm_limit={primary: 1})
responses: Final = burst(enforcing.gateway, tuple((limited, marker) for marker in markers), primary)
served: Final = served_once_in_a_burst(enforcing, markers, responses)
assert enforcing.traffic() == enforcing.expected({primary: (served,)})
def test_config_list_reports_the_setting_as_an_unset_boolean(gateway: Gateway) -> None:
listed: Final = listed_setting(gateway)
assert [(entry["field_type"], entry["field_value"], entry["stored_in_db"]) for entry in listed] == [
("Boolean", None, None)
], listed
@pytest.mark.parametrize(
"value", tuple(value for _, value in MALFORMED_VALUES), ids=tuple(name for name, _ in MALFORMED_VALUES)
)
def test_config_field_update_rejects_a_non_boolean_value_and_stores_nothing(gateway: Gateway, value: JsonValue) -> None:
response: Final = update_setting(gateway, value)
try:
assert response.status_code == 400, response.text
assert JSON_OBJECT.validate_json(response.content) == {
"detail": {"error": f"Invalid type of field value={type(value)} passed in."}
}, response.text
listed: Final = listed_setting(gateway)
assert [(entry["field_value"], entry["stored_in_db"]) for entry in listed] == [(None, None)], listed
finally:
if response.status_code == 200:
gateway.post("/config/field/delete", {"field_name": SETTING, "config_type": "general_settings"})

View file

@ -0,0 +1,194 @@
from __future__ import annotations
from collections.abc import Generator, Iterator, Mapping, Sequence
from contextlib import ExitStack, contextmanager
from pathlib import Path
from types import MappingProxyType
from typing import Final
import httpx
import pytest
from integration._support.client import eventually, gateway_from_environment, string_value
from integration._support.database import scratch_database
from integration._support.process import graceful_stop_seconds, owned_proxy_process
from integration._support.wire import Wire, wire_server
from integration.routing.test_per_model_rate_limit_fallbacks import (
SETTING,
Groups,
Rig,
assert_refused,
assert_served,
chat,
error_message,
healthy_reply,
listed_setting,
model_entry,
new_groups,
new_marker,
proxy_config,
update_setting,
)
from pydantic import JsonValue
pytestmark: Final = pytest.mark.timeout(4 * graceful_stop_seconds() + 120)
NO_OVERRIDES: Final[Mapping[str, str]] = MappingProxyType({})
LEGACY_LIMITER: Final = MappingProxyType({"LEGACY_MULTI_INSTANCE_RATE_LIMITING": "true"})
LEGACY_ATTEMPTS: Final = 10
NOT_A_BOOLEAN: Final = "is not a boolean, treating it as disabled"
def serving_groups(groups: Groups) -> tuple[str, ...]:
return (groups.primary, groups.fallback, groups.last)
def fallback_config(
directory: Path, groups: Groups, wires: Mapping[str, Wire], general_settings: Mapping[str, JsonValue]
) -> Path:
return proxy_config(
directory,
model_list=tuple(model_entry(group, wires[group]) for group in serving_groups(groups)),
fallbacks=({groups.primary: [groups.fallback]},),
general_settings=general_settings,
)
@contextmanager
def fallback_rig(
directory: Path, general_settings: Mapping[str, JsonValue], environment: Mapping[str, str] = NO_OVERRIDES
) -> Generator[Rig]:
groups: Final = new_groups()
with ExitStack() as stack:
rig_gateway: Final = stack.enter_context(gateway_from_environment())
wires: Final = MappingProxyType(
{group: stack.enter_context(wire_server(healthy_reply)) for group in serving_groups(groups)}
)
config: Final = fallback_config(directory, groups, wires, general_settings)
owned: Final = stack.enter_context(
owned_proxy_process(rig_gateway, directory, environment, config=config, workers=2)
)
yield Rig(owned.gateway, owned.log, groups, wires)
@pytest.fixture(scope="module")
def lenient(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
with fallback_rig(tmp_path_factory.mktemp("per-model-limits-lenient"), {SETTING: None}) as rig:
yield rig
def through_first_refusal(responses: Iterator[httpx.Response]) -> Iterator[httpx.Response]:
for response in responses:
yield response
if response.status_code != 200:
return
def sent_until_refused(rig: Rig, key: str, markers: Sequence[str]) -> tuple[httpx.Response, ...]:
return tuple(through_first_refusal(chat(rig.gateway, key, rig.groups.primary, marker) for marker in markers))
def test_null_setting_keeps_falling_back_on_a_per_model_429_without_a_warning(lenient: Rig) -> None:
lenient.traffic()
primary, fallback = lenient.groups.primary, lenient.groups.fallback
first, second = new_marker(), new_marker()
with lenient.gateway.scenario() as scenario:
key: Final = scenario.key(model_rpm_limit={primary: 1})
assert_served(chat(lenient.gateway, key, primary, first), primary, f"chatcmpl-{first}")
assert_served(chat(lenient.gateway, key, primary, second), fallback, f"chatcmpl-{second}")
assert lenient.traffic() == lenient.expected({primary: (first,), fallback: (second,)})
assert NOT_A_BOOLEAN not in lenient.log.read_text(errors="replace")
def test_key_router_settings_fallbacks_still_serve_a_per_model_429_when_the_setting_is_off(lenient: Rig) -> None:
lenient.traffic()
primary, last = lenient.groups.primary, lenient.groups.last
first, second = new_marker(), new_marker()
with lenient.gateway.scenario() as scenario:
key: Final = scenario.key(model_rpm_limit={primary: 1}, router_settings={"fallbacks": [{primary: [last]}]})
assert_served(chat(lenient.gateway, key, primary, first), primary, f"chatcmpl-{first}")
assert_served(chat(lenient.gateway, key, primary, second), last, f"chatcmpl-{second}")
assert lenient.traffic() == lenient.expected({primary: (first,), last: (second,)})
def test_request_disable_fallbacks_answers_a_per_model_429_when_the_setting_is_off(lenient: Rig) -> None:
lenient.traffic()
primary: Final = lenient.groups.primary
first, second = new_marker(), new_marker()
with lenient.gateway.scenario() as scenario:
key: Final = scenario.key(model_rpm_limit={primary: 1})
assert_served(chat(lenient.gateway, key, primary, first), primary, f"chatcmpl-{first}")
assert_refused(chat(lenient.gateway, key, primary, second, disable_fallbacks=True), "model_per_key", primary)
assert lenient.traffic() == lenient.expected({primary: (first,)})
def test_yaml_true_string_turns_the_setting_on(tmp_path: Path) -> None:
with fallback_rig(tmp_path, {SETTING: "true"}) as rig, rig.gateway.scenario() as scenario:
primary: Final = rig.groups.primary
first, second = new_marker(), new_marker()
key: Final = scenario.key(model_rpm_limit={primary: 1})
assert_served(chat(rig.gateway, key, primary, first), primary, f"chatcmpl-{first}")
assert_refused(chat(rig.gateway, key, primary, second), "model_per_key", primary)
assert rig.traffic() == rig.expected({primary: (first,)})
@pytest.mark.parametrize("raw", ("sometimes", ""), ids=("word", "empty"))
def test_yaml_non_boolean_setting_warns_and_keeps_falling_back(tmp_path: Path, raw: str) -> None:
with fallback_rig(tmp_path, {SETTING: raw}) as rig, rig.gateway.scenario() as scenario:
primary, fallback = rig.groups.primary, rig.groups.fallback
first, second = new_marker(), new_marker()
key: Final = scenario.key(model_rpm_limit={primary: 1})
assert_served(chat(rig.gateway, key, primary, first), primary, f"chatcmpl-{first}")
assert_served(chat(rig.gateway, key, primary, second), fallback, f"chatcmpl-{second}")
assert rig.traffic() == rig.expected({primary: (first,), fallback: (second,)})
eventually(
lambda: rig.log.read_text(errors="replace"),
lambda text: f"general_settings.{SETTING}={raw!r} {NOT_A_BOOLEAN}" in text,
seconds=10,
)
def test_legacy_limiter_per_model_429_skips_fallbacks(tmp_path: Path) -> None:
with fallback_rig(tmp_path, {SETTING: True}, LEGACY_LIMITER) as rig, rig.gateway.scenario() as scenario:
primary: Final = rig.groups.primary
markers: Final = tuple(new_marker() for _ in range(LEGACY_ATTEMPTS))
responses: Final = sent_until_refused(rig, scenario.key(model_rpm_limit={primary: 1}), markers)
*served, refused = responses
assert refused.status_code == 429, tuple(
(response.status_code, response.headers.get("x-litellm-model-group")) for response in responses
)
assert "LiteLLM Rate Limit Handler for rate limit type = model_per_key." in error_message(refused.text)
for marker, response in zip(markers, served):
assert_served(response, primary, f"chatcmpl-{marker}")
assert rig.traffic() == rig.expected({primary: markers[: len(served)]})
def test_database_setting_survives_a_restart_and_refuses_per_model_overflow(tmp_path: Path) -> None:
groups: Final = new_groups()
with ExitStack() as stack:
rig_gateway: Final = stack.enter_context(gateway_from_environment())
database_url: Final = stack.enter_context(scratch_database())
wires: Final = MappingProxyType(
{group: stack.enter_context(wire_server(healthy_reply)) for group in serving_groups(groups)}
)
config: Final = fallback_config(tmp_path, groups, wires, {})
environment: Final = MappingProxyType({"DATABASE_URL": database_url})
replica: Final = ("DATABASE_URL_READ_REPLICA",)
with owned_proxy_process(
rig_gateway, tmp_path, environment, config=config, remove_environment=replica, workers=2
) as writer:
writes: Final = (update_setting(writer.gateway, True), update_setting(writer.gateway, True))
assert tuple(write.status_code for write in writes) == (200, 200), tuple(write.text for write in writes)
restarted: Final = stack.enter_context(
owned_proxy_process(
rig_gateway, tmp_path, environment, config=config, remove_environment=replica, workers=2
)
)
listed: Final = listed_setting(restarted.gateway)
assert [(entry["field_value"], entry["stored_in_db"]) for entry in listed] == [(True, True)], listed
rig: Final = Rig(restarted.gateway, restarted.log, groups, wires)
rig.traffic()
key: Final = string_value(rig.gateway.post("/key/generate", {"model_rpm_limit": {groups.primary: 1}})["key"])
first, second = new_marker(), new_marker()
assert_served(chat(rig.gateway, key, groups.primary, first), groups.primary, f"chatcmpl-{first}")
assert_refused(chat(rig.gateway, key, groups.primary, second), "model_per_key", groups.primary)
assert rig.traffic() == rig.expected({groups.primary: (first,)})

View file

@ -3,7 +3,7 @@ Unit Tests for the max parallel request limiter v1 for the proxy
"""
import itertools
from collections.abc import Callable, Iterator
from collections.abc import Callable, Iterator, Mapping
from datetime import datetime
from typing import Final
@ -40,6 +40,39 @@ def _clock_rolling_over_after_first_read() -> Callable[[], datetime]:
)
@pytest.mark.parametrize(
"current,rpm_limit",
[
(None, 0),
({"current_requests": 0, "current_tpm": 0, "current_rpm": 1}, 1),
],
)
@pytest.mark.asyncio
async def test_model_per_key_rate_limit_error_carries_descriptor_key(
current: Mapping[str, int] | None, rpm_limit: int
):
handler: Final = PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
with pytest.raises(ProxyRateLimitError) as exc_info:
await handler.check_key_in_limits(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data={"model": "gpt-4o-mini"},
call_type="completion",
max_parallel_requests=10,
tpm_limit=100,
rpm_limit=rpm_limit,
current=dict(current) if current is not None else None,
request_count_api_key="test-key:model_per_key",
rate_limit_type="model_per_key",
values_to_update_in_cache=[],
)
assert exc_info.value.descriptor_key == "model_per_key"
@pytest.mark.asyncio
async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token():
handler = PROXY_MaxParallelRequestsHandler(

View file

@ -7467,6 +7467,36 @@ def test_rate_limit_error_reports_reset_time_in_utc_on_a_non_utc_proxy(process_t
)
def test_per_model_rate_limit_error_carries_descriptor_key() -> None:
handler: Final = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
over_limit: Final[RateLimitResponse] = {
"overall_code": "OVER_LIMIT",
"statuses": [
{
"code": "OVER_LIMIT",
"descriptor_key": "model_per_key",
"descriptor_value": "gpt-4o-mini",
"limit_remaining": 0,
"rate_limit_type": "requests",
"current_limit": 2,
}
],
}
with pytest.raises(ProxyRateLimitError) as exc_info:
handler._handle_rate_limit_error(
response=over_limit,
descriptors=[
{"key": "model_per_key", "value": "gpt-4o-mini", "rate_limit": None}
],
requested_model="gpt-4o-mini",
)
assert exc_info.value.descriptor_key == "model_per_key"
def _resolve_alias_to_target(model: str) -> str | None:
return "target" if model == "alias" else None

View file

@ -2,6 +2,7 @@ import asyncio
import copy
import datetime
import json
from collections.abc import Mapping
from types import MappingProxyType, SimpleNamespace
from typing import AsyncGenerator, Callable, Final, Iterator, Literal, Optional, Sequence
from urllib.parse import unquote_plus
@ -7095,12 +7096,16 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
user_api_key_dict: ProxyUserAPIKeyAuth,
fallbacks: list[dict[str, list[str]]],
model_guardrails: dict[str, list[str]] | None = None,
untagged_rate_limited_models: frozenset[str] = frozenset(),
) -> tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]]:
"""Real v3 limiter (the default ``parallel_request_limiter``) wired in through the
``proxy_logging_obj`` seam, so ``common_processing_pre_call_logic`` runs for real:
``add_litellm_data_to_request`` with a live OTel span, ``function_setup``, then the limiter."""
``add_litellm_data_to_request`` with a live OTel span, ``function_setup``, then the limiter.
``untagged_rate_limited_models`` are rejected before the limiter with a 429 that carries no
descriptor, the way the dynamic and batch limiters raise."""
from litellm.caching.caching import DualCache
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3
from litellm.proxy.utils import InternalUsageCache
@ -7115,6 +7120,11 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
skip_guardrails: bool = False,
) -> dict[str, object]:
limiter_models.append(str(data["model"]))
if data["model"] in untagged_rate_limited_models:
raise ProxyRateLimitError(
detail=f"Priority rate limit exceeded for {data['model']}",
headers={"retry-after": "1"},
)
await limiter.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
@ -7146,37 +7156,49 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
@staticmethod
def _otel_key(
rpm_limit: int | None = None,
api_key: str = "hashed-key",
model_rpm_limit: dict[str, int] | None = None,
disable_fallbacks: bool | None = None,
team_model_rpm_limit: dict[str, int] | None = None,
organization_model_rpm_limit: dict[str, int] | None = None,
project_metadata: Mapping[str, Mapping[str, int]] | None = None,
) -> ProxyUserAPIKeyAuth:
from opentelemetry.sdk.trace import TracerProvider
span = TracerProvider().get_tracer("test").start_span("proxy-request")
span: Final = TracerProvider().get_tracer("test").start_span("proxy-request")
return ProxyUserAPIKeyAuth(
api_key="hashed-key",
api_key=api_key,
parent_otel_span=span,
rpm_limit=rpm_limit,
metadata={
**({"model_rpm_limit": model_rpm_limit} if model_rpm_limit else {}),
**({"disable_fallbacks": disable_fallbacks} if disable_fallbacks is not None else {}),
},
team_id="team-1" if team_model_rpm_limit else None,
team_metadata={"model_rpm_limit": team_model_rpm_limit} if team_model_rpm_limit else None,
org_id="org-1" if organization_model_rpm_limit else None,
organization_metadata=(
{"model_rpm_limit": organization_model_rpm_limit} if organization_model_rpm_limit else None
),
project_id="project-1" if project_metadata else None,
project_metadata={k: dict(v) for k, v in project_metadata.items()} if project_metadata else None,
)
@staticmethod
def _chat_request() -> Request:
return Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []})
async def _pre_call(
async def _run(
self,
data: dict[str, object],
processor: ProxyBaseLLMRequestProcessing,
user_api_key_dict: ProxyUserAPIKeyAuth,
rig: tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]],
) -> tuple[ProxyBaseLLMRequestProcessing, tuple[dict[str, object], LiteLLMLoggingObj]]:
general_settings: Mapping[str, object] | None = None,
) -> tuple[dict[str, object], LiteLLMLoggingObj]:
proxy_logging_obj, router, proxy_config, _ = rig
processor = ProxyBaseLLMRequestProcessing(data=data)
result = await processor._pre_call_with_fallbacks(
return await processor._pre_call_with_fallbacks(
request=self._chat_request(),
general_settings={},
general_settings=dict(general_settings or {}),
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
version=None,
@ -7190,7 +7212,16 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
route_type="acompletion",
llm_router=router,
)
return processor, result
async def _pre_call(
self,
data: dict[str, object],
user_api_key_dict: ProxyUserAPIKeyAuth,
rig: tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]],
general_settings: Mapping[str, object] | None = None,
) -> tuple[ProxyBaseLLMRequestProcessing, tuple[dict[str, object], LiteLLMLoggingObj]]:
processor = ProxyBaseLLMRequestProcessing(data=data)
return processor, await self._run(processor, user_api_key_dict, rig, general_settings)
@pytest.mark.asyncio
async def test_v3_limiter_with_otel_span_falls_back_from_client_request(self, monkeypatch: pytest.MonkeyPatch):
@ -7265,6 +7296,232 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
assert processor.data["litellm_logging_obj"].model == primary_model
assert processor.data["litellm_call_id"]
HARD_PER_MODEL_LIMITS: Final = MappingProxyType({"disable_fallbacks_on_per_model_rate_limits": True})
@pytest.mark.parametrize("cap_owner", ["key", "team"])
@pytest.mark.asyncio
async def test_v3_limiter_per_model_cap_returns_429_when_per_model_limits_are_hard(
self, cap_owner: str, monkeypatch: pytest.MonkeyPatch
):
"""A key's own per-model cap and one inherited from its team both resolve into the key's
per-model descriptor (see ``get_key_model_rpm_limit``), so both must stop the fallback."""
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
primary_model: Final = "gpt-4.1"
fallback_model: Final = "gpt-4.1-mini"
key: Final = (
self._otel_key(model_rpm_limit={primary_model: 1})
if cap_owner == "key"
else self._otel_key(team_model_rpm_limit={primary_model: 1})
)
rig: Final = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request: Final = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
_, (first_data, _) = await self._pre_call(dict(request), key, rig, self.HARD_PER_MODEL_LIMITS)
processor: Final = ProxyBaseLLMRequestProcessing(data=dict(request))
with pytest.raises(ProxyRateLimitError) as exc_info:
await self._run(processor, key, rig, self.HARD_PER_MODEL_LIMITS)
assert first_data["model"] == primary_model
assert rig[3] == [primary_model, primary_model]
assert exc_info.value.status_code == 429
assert exc_info.value.descriptor_key == "model_per_key"
assert exc_info.value.headers["retry-after"]
assert processor.data["model"] == primary_model
@pytest.mark.asyncio
async def test_v3_limiter_team_cap_shared_by_two_keys_returns_429_when_per_model_limits_are_hard(
self, monkeypatch: pytest.MonkeyPatch
):
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
primary_model: Final = "gpt-4.1"
fallback_model: Final = "gpt-4.1-mini"
first_key: Final = self._otel_key(team_model_rpm_limit={primary_model: 1})
second_key: Final = self._otel_key(api_key="hashed-key-2", team_model_rpm_limit={primary_model: 1})
rig: Final = self._v3_limiter_rig(monkeypatch, first_key, [{primary_model: [fallback_model]}])
request: Final = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
_, (first_data, _) = await self._pre_call(dict(request), first_key, rig, self.HARD_PER_MODEL_LIMITS)
processor: Final = ProxyBaseLLMRequestProcessing(data=dict(request))
with pytest.raises(ProxyRateLimitError) as exc_info:
await self._run(processor, second_key, rig, self.HARD_PER_MODEL_LIMITS)
assert first_data["model"] == primary_model
assert rig[3] == [primary_model, primary_model]
assert exc_info.value.status_code == 429
assert exc_info.value.descriptor_key == "model_per_team"
assert exc_info.value.headers["retry-after"]
assert processor.data["model"] == primary_model
@pytest.mark.parametrize("setting_value", [True, "true", "True", "1"])
@pytest.mark.asyncio
async def test_per_model_limits_are_hard_for_true_and_a_true_string(
self, setting_value: object, monkeypatch: pytest.MonkeyPatch
):
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
primary_model: Final = "gpt-4.1"
fallback_model: Final = "gpt-4.1-mini"
key: Final = self._otel_key(model_rpm_limit={primary_model: 1})
rig: Final = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request: Final = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
general_settings: Final = {"disable_fallbacks_on_per_model_rate_limits": setting_value}
await self._pre_call(dict(request), key, rig, general_settings)
processor: Final = ProxyBaseLLMRequestProcessing(data=dict(request))
with pytest.raises(ProxyRateLimitError) as exc_info:
await self._run(processor, key, rig, general_settings)
assert exc_info.value.descriptor_key == "model_per_key"
assert rig[3] == [primary_model, primary_model]
assert processor.data["model"] == primary_model
@pytest.mark.parametrize("setting_value", [False, "false", "", None, "not-a-bool"])
@pytest.mark.asyncio
async def test_per_model_limits_stay_soft_for_every_other_setting_value(
self, setting_value: object, monkeypatch: pytest.MonkeyPatch
):
primary_model: Final = "gpt-4.1"
fallback_model: Final = "gpt-4.1-mini"
key: Final = self._otel_key(model_rpm_limit={primary_model: 1})
rig: Final = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request: Final = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
general_settings: Final = {"disable_fallbacks_on_per_model_rate_limits": setting_value}
await self._pre_call(dict(request), key, rig, general_settings)
_, (data, _) = await self._pre_call(dict(request), key, rig, general_settings)
assert data["model"] == fallback_model
assert rig[3] == [primary_model, primary_model, fallback_model]
@pytest.mark.parametrize(
("cap_owner", "expected_descriptor"),
[
("organization", "model_per_organization"),
("project", "model_per_project"),
("project_itpm", "model_per_project_itpm"),
("project_otpm", "model_per_project_otpm"),
],
)
@pytest.mark.asyncio
async def test_v3_limiter_org_and_project_caps_return_429_when_per_model_limits_are_hard(
self, cap_owner: str, expected_descriptor: str, monkeypatch: pytest.MonkeyPatch
):
"""Every per-model descriptor the real limiter raises stops the fallback hunt, each under its
own key: an organization cap, a project cap, and the project's input and output token caps
(the token caps trip on the first request, since its own tokens already exceed a limit of 1)."""
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
primary_model: Final = "gpt-4.1"
fallback_model: Final = "gpt-4.1-mini"
key: Final = (
self._otel_key(organization_model_rpm_limit={primary_model: 1})
if cap_owner == "organization"
else self._otel_key(
project_metadata={
{
"project": "model_rpm_limit",
"project_itpm": "model_itpm_limit",
"project_otpm": "model_otpm_limit",
}[cap_owner]: {primary_model: 1}
}
)
)
rig: Final = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request: Final = {
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 8,
}
processors: Final = tuple(ProxyBaseLLMRequestProcessing(data=dict(request)) for _ in range(2))
async def run_both_requests() -> None:
for processor in processors:
await self._run(processor, key, rig, self.HARD_PER_MODEL_LIMITS)
with pytest.raises(ProxyRateLimitError) as exc_info:
await run_both_requests()
assert set(rig[3]) == {primary_model}
assert exc_info.value.status_code == 429
assert exc_info.value.descriptor_key == expected_descriptor
assert exc_info.value.headers["retry-after"]
assert all(processor.data["model"] == primary_model for processor in processors)
@pytest.mark.asyncio
async def test_v3_limiter_global_key_cap_still_tries_fallbacks_when_per_model_limits_are_hard(
self, monkeypatch: pytest.MonkeyPatch
):
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
primary_model: Final = "gpt-4.1"
fallback_model: Final = "gpt-4.1-mini"
key: Final = self._otel_key(rpm_limit=1)
rig: Final = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request: Final = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
await self._pre_call(dict(request), key, rig, self.HARD_PER_MODEL_LIMITS)
processor: Final = ProxyBaseLLMRequestProcessing(data=dict(request))
with pytest.raises(ProxyRateLimitError) as exc_info:
await self._run(processor, key, rig, self.HARD_PER_MODEL_LIMITS)
assert rig[3] == [primary_model, primary_model, fallback_model]
assert exc_info.value.descriptor_key == "api_key"
assert processor.data["model"] == primary_model
@pytest.mark.asyncio
async def test_v3_limiter_fallback_per_model_cap_returns_429_when_per_model_limits_are_hard(
self, monkeypatch: pytest.MonkeyPatch
):
"""An untagged 429 on the primary still opens the fallback hunt, but a per-model cap on the
first fallback ends it with that fallback's 429 instead of moving on to the next model."""
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
primary_model: Final = "gpt-4.1"
capped_fallback: Final = "gpt-4.1-mini"
last_fallback: Final = "gpt-4.1-nano"
key: Final = self._otel_key(model_rpm_limit={capped_fallback: 1})
rig: Final = self._v3_limiter_rig(
monkeypatch,
key,
[{primary_model: [capped_fallback, last_fallback]}],
untagged_rate_limited_models=frozenset({primary_model}),
)
request: Final = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
_, (first_data, _) = await self._pre_call(dict(request), key, rig, self.HARD_PER_MODEL_LIMITS)
processor: Final = ProxyBaseLLMRequestProcessing(data=dict(request))
with pytest.raises(ProxyRateLimitError) as exc_info:
await self._run(processor, key, rig, self.HARD_PER_MODEL_LIMITS)
assert first_data["model"] == capped_fallback
assert rig[3] == [primary_model, capped_fallback, primary_model, capped_fallback]
assert exc_info.value.descriptor_key == "model_per_key"
assert processor.data["model"] == primary_model
@pytest.mark.asyncio
async def test_v3_limiter_fallback_per_model_cap_moves_on_when_setting_is_off(
self, monkeypatch: pytest.MonkeyPatch
):
primary_model: Final = "gpt-4.1"
capped_fallback: Final = "gpt-4.1-mini"
last_fallback: Final = "gpt-4.1-nano"
key: Final = self._otel_key(model_rpm_limit={capped_fallback: 1})
rig: Final = self._v3_limiter_rig(
monkeypatch,
key,
[{primary_model: [capped_fallback, last_fallback]}],
untagged_rate_limited_models=frozenset({primary_model}),
)
request: Final = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
await self._pre_call(dict(request), key, rig)
_, (data, _) = await self._pre_call(dict(request), key, rig)
assert data["model"] == last_fallback
assert rig[3] == [primary_model, capped_fallback, primary_model, capped_fallback, last_fallback]
@pytest.mark.asyncio
async def test_fallback_lookup_uses_alias_resolved_model_group(self, monkeypatch: pytest.MonkeyPatch):
primary_model = "gpt-4.1"

View file

@ -29438,6 +29438,11 @@ export interface components {
* @description If True, disables signing in to the Admin UI with the environment credentials: UI_USERNAME/UI_PASSWORD, or the master key when UI_PASSWORD is unset (that fallback means env-credential login is always live by default). Database users with passwords are unaffected. LOCKOUT RISK: create at least one proxy admin user with a password before enabling, or nobody can sign in to the UI. A locked-out admin can still administer the proxy over the API with the master key, and can unset this setting and restart the proxy to restore env-credential login. Default is False.
*/
disable_env_credential_login?: boolean | null;
/**
* Disable Fallbacks On Per Model Rate Limits
* @description If true, a request rejected by a key/team/org/project per-model rate limit (model_rpm_limit / model_tpm_limit) returns 429 instead of retrying on the configured fallbacks
*/
disable_fallbacks_on_per_model_rate_limits?: boolean | null;
/**
* Disable Password Login When Sso Enabled
* @description If True and SSO is configured (MICROSOFT_CLIENT_ID, GOOGLE_CLIENT_ID, GENERIC_CLIENT_ID, or SAML_IDP_METADATA_URL/XML), disables username/password login on /login, /v2/login, and /v3/login so SSO is the only way to reach the Admin UI. An admin locked out of the UI can still administer the proxy over the API with the master key; unset this setting and restart the proxy to restore UI username/password login. Default is False.