From 15154a40a83fe2fdf6ca9e8f5b3ba438056a5bd9 Mon Sep 17 00:00:00 2001 From: Shivam Rawat Date: Fri, 9 Oct 2026 23:10:23 -0700 Subject: [PATCH] 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> --- litellm/proxy/_types.py | 8 + litellm/proxy/common_request_processing.py | 22 +- .../common_utils/proxy_rate_limit_error.py | 31 + .../proxy/hooks/parallel_request_limiter.py | 4 + .../hooks/parallel_request_limiter_v3.py | 1 + litellm/proxy/proxy_server.py | 1 + .../test_per_model_rate_limit_fallbacks.py | 647 ++++++++++++++++++ ...er_model_rate_limit_fallbacks_lifecycle.py | 194 ++++++ .../hooks/test_parallel_request_limiter.py | 35 +- .../hooks/test_parallel_request_limiter_v3.py | 30 + .../proxy/test_common_request_processing.py | 277 +++++++- ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 + 12 files changed, 1242 insertions(+), 13 deletions(-) create mode 100644 tests/integration/routing/test_per_model_rate_limit_fallbacks.py create mode 100644 tests/integration/routing/test_per_model_rate_limit_fallbacks_lifecycle.py diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ff7f071d534..477edf47c5f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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=( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 4f18c773a24..7f0e0b1fc87 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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 diff --git a/litellm/proxy/common_utils/proxy_rate_limit_error.py b/litellm/proxy/common_utils/proxy_rate_limit_error.py index b028e7fda20..089d0b7a226 100644 --- a/litellm/proxy/common_utils/proxy_rate_limit_error.py +++ b/litellm/proxy/common_utils/proxy_rate_limit_error.py @@ -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 diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 7c02ac308f3..d3ca4597f0d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -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") diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 5b865f3b5c3..21025a8a820 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a7198fd1a95..dac46855fd6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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", diff --git a/tests/integration/routing/test_per_model_rate_limit_fallbacks.py b/tests/integration/routing/test_per_model_rate_limit_fallbacks.py new file mode 100644 index 00000000000..428a485f397 --- /dev/null +++ b/tests/integration/routing/test_per_model_rate_limit_fallbacks.py @@ -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"}) diff --git a/tests/integration/routing/test_per_model_rate_limit_fallbacks_lifecycle.py b/tests/integration/routing/test_per_model_rate_limit_fallbacks_lifecycle.py new file mode 100644 index 00000000000..4361843b2c9 --- /dev/null +++ b/tests/integration/routing/test_per_model_rate_limit_fallbacks_lifecycle.py @@ -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,)}) diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter.py b/tests/unit/proxy/hooks/test_parallel_request_limiter.py index a2e43b3bc9c..89cb6bd1ef2 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter.py @@ -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( diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index d9cf8536912..55163c6240b 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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 diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index d2065d34bd3..b040cc861bf 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -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" diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 6a34399236c..c9bf2e4aa5d 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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.