From 6be9c4a978c2a35d9cae254e2c05a1dc536d2fd4 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 03:17:49 +0000 Subject: [PATCH] refactor(router): compose DeploymentSemaphore over asyncio.Semaphore instead of subclassing it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/router.py | 6 ++-- .../client_initalization_utils.py | 31 ++++++++++++++----- tests/test_litellm/test_router.py | 5 ++- 3 files changed, 30 insertions(+), 12 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index b62c83b8ab1..132f48730df 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -148,7 +148,7 @@ from litellm.router_utils.batch_utils import ( replace_model_in_jsonl, should_replace_model_in_jsonl, ) -from litellm.router_utils.client_initalization_utils import InitalizeCachedClient +from litellm.router_utils.client_initalization_utils import DeploymentSemaphore, InitalizeCachedClient from litellm.router_utils.clientside_credential_handler import ( get_dynamic_litellm_params, is_clientside_credential, @@ -3640,7 +3640,7 @@ class Router: client_type="max_parallel_requests", ) async with contextlib.AsyncExitStack() as deployment_slot: - if isinstance(rpm_semaphore, asyncio.Semaphore): + if isinstance(rpm_semaphore, DeploymentSemaphore): await deployment_slot.enter_async_context(rpm_semaphore) await self.async_routing_strategy_pre_call_checks( deployment=deployment, @@ -8512,7 +8512,7 @@ class Router: client_type="max_parallel_requests", ) async with contextlib.AsyncExitStack() as slot: - if isinstance(rpm_semaphore, asyncio.Semaphore): + if isinstance(rpm_semaphore, DeploymentSemaphore): await slot.enter_async_context(rpm_semaphore) await self.async_routing_strategy_pre_call_checks(deployment=deployment, parent_otel_span=parent_otel_span) yield diff --git a/litellm/router_utils/client_initalization_utils.py b/litellm/router_utils/client_initalization_utils.py index a135978d09e..72854cba028 100644 --- a/litellm/router_utils/client_initalization_utils.py +++ b/litellm/router_utils/client_initalization_utils.py @@ -1,5 +1,6 @@ import asyncio import time +from types import TracebackType from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_router_logger @@ -15,22 +16,36 @@ else: LitellmRouter = Any -class DeploymentSemaphore(asyncio.Semaphore): +class DeploymentSemaphore: """A deployment's max_parallel_requests slots. ``queue_size=None`` parks callers without bound, like a plain ``asyncio.Semaphore``; otherwise a caller arriving while all slots are busy and ``queue_size`` callers already wait gets a 429 instead of being parked.""" def __init__(self, max_parallel_requests: int, model_id: str, model_group: str, queue_size: int | None) -> None: - super().__init__(max_parallel_requests) - self.max_parallel_requests = max_parallel_requests - self.model_id = model_id - self.model_group = model_group + self._slots: Final = asyncio.Semaphore(max_parallel_requests) + self.max_parallel_requests: Final = max_parallel_requests + self.model_id: Final = model_id + self.model_group: Final = model_group self.queue_size = queue_size self.waiting = 0 + def locked(self) -> bool: + return self._slots.locked() + + def release(self) -> None: + self._slots.release() + + async def __aenter__(self) -> None: + await self.acquire() + + async def __aexit__( + self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: TracebackType | None + ) -> None: + self._slots.release() + async def acquire(self) -> bool: - if not self.locked(): - return await super().acquire() + if not self._slots.locked(): + return await self._slots.acquire() if self.queue_size is not None and self.waiting >= self.queue_size: raise RateLimitError( message=( @@ -57,7 +72,7 @@ class DeploymentSemaphore(asyncio.Semaphore): self.queue_size, ) try: - return await super().acquire() + return await self._slots.acquire() finally: self.waiting -= 1 verbose_router_logger.debug( diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index b094312808e..a674f767dde 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -47,6 +47,7 @@ from litellm.router import ( _is_retriable_anthropic_status, ) from litellm.router_strategy import simple_shuffle +from litellm.router_utils.client_initalization_utils import DeploymentSemaphore from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy @@ -1520,7 +1521,9 @@ async def test_router_ageneric_api_call_with_fallbacks_helper(): }, } - mock_semaphore = asyncio.Semaphore(1) + mock_semaphore = DeploymentSemaphore( + max_parallel_requests=1, model_id="deployment-1", model_group="gpt-3.5-turbo", queue_size=None + ) with patch.object( router, "_update_kwargs_with_deployment"