Skip to content
446 changes: 116 additions & 330 deletions litellm/router.py

Large diffs are not rendered by default.

49 changes: 46 additions & 3 deletions litellm/router_utils/client_initalization_utils.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
import asyncio
from types import TracebackType
from typing import TYPE_CHECKING, Any, Final

from litellm.exceptions import RateLimitError, RateLimitErrorCategory, RateLimitType
from litellm.types.router import RouterErrors
from litellm.utils import calculate_max_parallel_requests

if TYPE_CHECKING:
Expand All @@ -11,6 +13,43 @@
LitellmRouter = Any


class MaxParallelRequestsLimit:
"""A deployment's max_parallel_requests slots. A caller arriving while every slot is in use gets a 429 instead
of waiting for one to free up."""

def __init__(self, max_parallel_requests: int, model_id: str, model_group: str) -> None:
self.max_parallel_requests: Final = max_parallel_requests
self.model_id: Final = model_id
self.model_group: Final = model_group
self.in_flight = 0

def __enter__(self) -> None:
self.acquire()

def __exit__(
self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: TracebackType | None
) -> None:
self.release()

def acquire(self) -> None:
if self.in_flight >= self.max_parallel_requests:
raise RateLimitError(
message=(
f"{RouterErrors.max_parallel_requests_exceeded.value} Deployment model_group={self.model_group}, "
f"id={self.model_id} already has max_parallel_requests={self.max_parallel_requests} requests in "
"flight. Raise max_parallel_requests (or the rpm/tpm it is derived from) for this deployment"
),
llm_provider="",
model=self.model_group,
category=RateLimitErrorCategory.LITELLM_RATE_LIMIT,
rate_limit_type=RateLimitType.CONCURRENT_REQUESTS,
Comment thread
greptile-apps[bot] marked this conversation as resolved.
)
self.in_flight += 1

def release(self) -> None:
self.in_flight -= 1


Comment thread
greptile-apps[bot] marked this conversation as resolved.
class InitalizeCachedClient:
@staticmethod
def set_max_parallel_requests_client(litellm_router_instance: LitellmRouter, model: dict):
Expand All @@ -26,10 +65,14 @@ def set_max_parallel_requests_client(litellm_router_instance: LitellmRouter, mod
default_max_parallel_requests=litellm_router_instance.default_max_parallel_requests,
)
if calculated_max_parallel_requests:
semaphore: Final = asyncio.Semaphore(calculated_max_parallel_requests)
limit: Final = MaxParallelRequestsLimit(
max_parallel_requests=calculated_max_parallel_requests,
model_id=model_id,
model_group=model.get("model_name", ""),
)
cache_key: Final = f"{model_id}_max_parallel_requests_client"
litellm_router_instance.cache.set_cache(
key=cache_key,
value=semaphore,
value=limit,
local_only=True,
)
1 change: 1 addition & 0 deletions litellm/types/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -653,6 +653,7 @@ class RouterErrors(enum.Enum):
"""

user_defined_ratelimit_error = "Deployment over user-defined ratelimit."
max_parallel_requests_exceeded = "Deployment has all max_parallel_requests slots in use."
no_deployments_available = "No deployments available for selected model"
all_deployments_in_cooldown = "All deployments for selected model are in cooldown"
no_deployments_with_tag_routing = "Not allowed to access model due to tags configuration"
Expand Down
13 changes: 7 additions & 6 deletions tests/local_testing/test_router_max_parallel_requests.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
from typing import Optional

import litellm
from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit
from litellm.utils import calculate_max_parallel_requests

"""
Expand Down Expand Up @@ -93,26 +94,26 @@ def test_setting_mpr_limits_per_model(
default_max_parallel_requests=default_max_parallel_requests,
)

mpr_client: Optional[asyncio.Semaphore] = router._get_client(
mpr_client: Optional[MaxParallelRequestsLimit] = router._get_client(
deployment=deployment,
kwargs={},
client_type="max_parallel_requests",
)

if max_parallel_requests is not None:
assert max_parallel_requests == mpr_client._value
assert max_parallel_requests == mpr_client.max_parallel_requests
elif rpm is not None:
assert rpm == mpr_client._value
assert rpm == mpr_client.max_parallel_requests
elif tpm is not None:
calculated_rpm = int(tpm / 1000 * 6)
if calculated_rpm == 0:
calculated_rpm = 1
print(
f"test calculated_rpm: {calculated_rpm}, calculated_max_parallel_requests={mpr_client._value}"
f"test calculated_rpm: {calculated_rpm}, calculated_max_parallel_requests={mpr_client.max_parallel_requests}"
)
assert calculated_rpm == mpr_client._value
assert calculated_rpm == mpr_client.max_parallel_requests
elif default_max_parallel_requests is not None:
assert mpr_client._value == default_max_parallel_requests
assert mpr_client.max_parallel_requests == default_max_parallel_requests
else:
assert mpr_client is None

Expand Down
125 changes: 125 additions & 0 deletions tests/test_litellm/router_utils/test_client_initalization_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
import asyncio
from typing import Final

import pytest

import litellm
from litellm import Router
from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit


def _limit(max_parallel_requests: int = 1) -> MaxParallelRequestsLimit:
return MaxParallelRequestsLimit(
max_parallel_requests=max_parallel_requests, model_id="deployment-1", model_group="gpt-5.6"
)


async def _hold(limit: MaxParallelRequestsLimit, release: asyncio.Event) -> str:
with limit:
await release.wait()
return "ok"


def _expect_rejection(limit: MaxParallelRequestsLimit) -> litellm.RateLimitError:
with pytest.raises(litellm.RateLimitError) as excinfo:
limit.acquire()
return excinfo.value


@pytest.mark.asyncio
async def test_request_arriving_while_every_slot_is_in_use_gets_429_without_waiting():
limit: Final = _limit(max_parallel_requests=2)
release: Final = asyncio.Event()
holders: Final = [asyncio.create_task(_hold(limit, release)) for _ in range(2)]
await asyncio.sleep(0)
assert limit.in_flight == 2

rejection: Final = _expect_rejection(limit)

assert rejection.status_code == 429
assert "deployment-1" in rejection.message
assert "gpt-5.6" in rejection.message
assert "max_parallel_requests=2" in rejection.message
assert limit.in_flight == 2

release.set()
assert await asyncio.wait_for(asyncio.gather(*holders), timeout=2) == ["ok", "ok"]
assert limit.in_flight == 0
with limit:
assert limit.in_flight == 1
assert limit.in_flight == 0


@pytest.mark.asyncio
async def test_burst_over_the_cap_admits_exactly_max_parallel_requests_and_rejects_the_rest():
limit: Final = _limit(max_parallel_requests=3)
release: Final = asyncio.Event()

async def attempt() -> str:
try:
return await _hold(limit, release)
except litellm.RateLimitError as e:
return f"rejected:{e.status_code}"

callers: Final = [asyncio.create_task(attempt()) for _ in range(10)]
await asyncio.sleep(0)
assert limit.in_flight == 3
release.set()
outcomes: Final = await asyncio.wait_for(asyncio.gather(*callers), timeout=2)
assert outcomes.count("ok") == 3
assert outcomes.count("rejected:429") == 7
assert limit.in_flight == 0


def test_slot_is_released_when_the_held_call_raises():
limit: Final = _limit()
with pytest.raises(RuntimeError):
with limit:
raise RuntimeError("provider blew up")
assert limit.in_flight == 0
with limit:
assert limit.in_flight == 1


def _router_limit(router: Router, model_name: str) -> MaxParallelRequestsLimit:
deployment: Final = router.get_deployment_by_model_group_name(model_group_name=model_name)
assert deployment is not None
client: Final = router._get_client(
deployment=deployment.model_dump(), kwargs={}, client_type="max_parallel_requests"
)
assert isinstance(client, MaxParallelRequestsLimit)
return client


@pytest.mark.parametrize(
("litellm_params", "expected_cap"),
[
({"max_parallel_requests": 2, "rpm": 7, "tpm": 100_000}, 2),
({"rpm": 7, "tpm": 100_000}, 7),
({"tpm": 100_000}, 600),
({"tpm": 100}, 1),
],
)
@pytest.mark.asyncio
async def test_router_deployment_rejects_past_its_derived_cap(litellm_params: dict[str, int], expected_cap: int):
router: Final = Router(
model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6", **litellm_params}}]
)
limit: Final = _router_limit(router, "gpt-5.6")
assert limit.max_parallel_requests == expected_cap
release: Final = asyncio.Event()
holders: Final = [asyncio.create_task(_hold(limit, release)) for _ in range(expected_cap)]
await asyncio.sleep(0)
assert limit.in_flight == expected_cap
assert f"max_parallel_requests={expected_cap}" in _expect_rejection(limit).message
release.set()
assert await asyncio.wait_for(asyncio.gather(*holders), timeout=2) == ["ok"] * expected_cap


def test_router_without_any_concurrency_setting_has_no_limit():
router: Final = Router(model_list=[{"model_name": "gpt-5.6", "litellm_params": {"model": "openai/gpt-5.6"}}])
deployment: Final = router.get_deployment_by_model_group_name(model_group_name="gpt-5.6")
assert deployment is not None
assert (
router._get_client(deployment=deployment.model_dump(), kwargs={}, client_type="max_parallel_requests") is None
)
Loading
Loading