diff --git a/litellm/router_strategy/least_busy.py b/litellm/router_strategy/least_busy.py index de6a4c4f59a..00d27b5f8d9 100644 --- a/litellm/router_strategy/least_busy.py +++ b/litellm/router_strategy/least_busy.py @@ -39,7 +39,7 @@ class _Deployment(TypedDict): _CALL_KWARGS: Final = TypeAdapter(_CallKwargs) _DEPLOYMENTS: Final = TypeAdapter(list[_Deployment]) _REDIS_COUNTS: Final = TypeAdapter(dict[str, float | None]) -_MEMORY_COUNTS: Final = TypeAdapter(tuple[float | None, ...]) +_MEMORY_COUNTS: Final = TypeAdapter(tuple[float | None, ...] | None) def _request_count_key(model_group: str, deployment_id: str) -> str: @@ -68,8 +68,20 @@ def _request_count_keys(model_group: str, healthy_deployments: Sequence[Mapping[ ) -def _as_count(value: float | None) -> int: - return 0 if value is None else int(value) +def _as_counts(values: Sequence[float | None]) -> tuple[int, ...]: + return tuple(0 if value is None else int(value) for value in values) + + +def _shared_counts(raw: object, keys: tuple[str, ...]) -> tuple[int, ...]: + by_key: Final = _REDIS_COUNTS.validate_python(raw) + return _as_counts([by_key.get(key) for key in keys]) + + +def _local_counts(raw: object, keys: tuple[str, ...]) -> tuple[int, ...]: + values: Final = _MEMORY_COUNTS.validate_python(raw) + if values is None or len(values) != len(keys): + return (0,) * len(keys) + return _as_counts(values) def _least_busy( @@ -82,7 +94,8 @@ def _least_busy( def _warn_unreadable(model_group: str, error: Exception) -> None: verbose_router_logger.warning( - "least-busy routing could not read the in-flight counts for %s, treating every deployment as idle: %s", + "least-busy routing could not read the shared in-flight counts for %s, " + "falling back to this worker's own counts: %s", model_group, error, ) @@ -135,23 +148,31 @@ class LeastBusyLoggingHandler(CustomLogger): self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]] ) -> Mapping[str, object] | None: keys: Final = _request_count_keys(model_group, healthy_deployments) - try: - counts: Final = tuple(_as_count(value) for value in self._read_counts(keys)) - except Exception as e: - _warn_unreadable(model_group, e) - return _least_busy(healthy_deployments, (0,) * len(keys)) - return _least_busy(healthy_deployments, counts) + redis_cache: Final = self.router_cache.redis_cache + if redis_cache is not None: + try: + shared: Final = _shared_counts(redis_cache.batch_get_cache(key_list=list(keys)), keys) + except Exception as e: + _warn_unreadable(model_group, e) + else: + return _least_busy(healthy_deployments, shared) + local: Final = _local_counts(self.router_cache.batch_get_cache(list(keys), local_only=True), keys) + return _least_busy(healthy_deployments, local) async def async_get_available_deployments( self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]] ) -> Mapping[str, object] | None: keys: Final = _request_count_keys(model_group, healthy_deployments) - try: - counts: Final = tuple(_as_count(value) for value in await self._async_read_counts(keys)) - except Exception as e: - _warn_unreadable(model_group, e) - return _least_busy(healthy_deployments, (0,) * len(keys)) - return _least_busy(healthy_deployments, counts) + redis_cache: Final = self.router_cache.redis_cache + if redis_cache is not None: + try: + shared: Final = _shared_counts(await redis_cache.async_batch_get_cache(key_list=list(keys)), keys) + except Exception as e: + _warn_unreadable(model_group, e) + else: + return _least_busy(healthy_deployments, shared) + local: Final = _local_counts(await self.router_cache.async_batch_get_cache(list(keys), local_only=True), keys) + return _least_busy(healthy_deployments, local) def _increment(self, kwargs: Mapping[str, object], delta: int) -> None: ref: Final = _deployment_ref(kwargs) @@ -160,10 +181,16 @@ class LeastBusyLoggingHandler(CustomLogger): key: Final = _request_count_key(*ref) redis_cache: Final = self.router_cache.redis_cache try: + local: Final = self.router_cache.increment_cache( + key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS + ) + if local < 0: + self.router_cache.set_cache(key, 0, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS) if redis_cache is None: - self.router_cache.increment_cache(key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS) - else: - redis_cache.increment_cache(key, delta, ttl=IN_FLIGHT_COUNT_TTL_SECONDS, refresh_ttl=True) + return + shared: Final = redis_cache.increment_cache(key, delta, ttl=IN_FLIGHT_COUNT_TTL_SECONDS, refresh_ttl=True) + if shared < 0: + redis_cache.set_cache(key, 0, ttl=IN_FLIGHT_COUNT_TTL_SECONDS) except Exception as e: _warn_unwritable(key, e) @@ -174,27 +201,17 @@ class LeastBusyLoggingHandler(CustomLogger): key: Final = _request_count_key(*ref) redis_cache: Final = self.router_cache.redis_cache try: + local: Final = await self.router_cache.async_increment_cache( + key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS + ) + if local is not None and local < 0: + await self.router_cache.async_set_cache(key, 0, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS) if redis_cache is None: - await self.router_cache.async_increment_cache( - key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS - ) - else: - await redis_cache.async_increment(key, delta, ttl=IN_FLIGHT_COUNT_TTL_SECONDS, refresh_ttl=True) + return + shared: Final = await redis_cache.async_increment( + key, delta, ttl=IN_FLIGHT_COUNT_TTL_SECONDS, refresh_ttl=True + ) + if shared < 0: + await redis_cache.async_set_cache(key, 0, ttl=IN_FLIGHT_COUNT_TTL_SECONDS) except Exception as e: _warn_unwritable(key, e) - - def _read_counts(self, keys: tuple[str, ...]) -> tuple[float | None, ...]: - redis_cache: Final = self.router_cache.redis_cache - if redis_cache is None: - return _MEMORY_COUNTS.validate_python(self.router_cache.batch_get_cache(list(keys), local_only=True)) - by_key: Final = _REDIS_COUNTS.validate_python(redis_cache.batch_get_cache(key_list=list(keys))) - return tuple(by_key.get(key) for key in keys) - - async def _async_read_counts(self, keys: tuple[str, ...]) -> tuple[float | None, ...]: - redis_cache: Final = self.router_cache.redis_cache - if redis_cache is None: - return _MEMORY_COUNTS.validate_python( - await self.router_cache.async_batch_get_cache(list(keys), local_only=True) - ) - by_key: Final = _REDIS_COUNTS.validate_python(await redis_cache.async_batch_get_cache(key_list=list(keys))) - return tuple(by_key.get(key) for key in keys) diff --git a/tests/test_litellm/router_strategy/test_least_busy.py b/tests/test_litellm/router_strategy/test_least_busy.py index 32a226a5d5b..55702c6fe73 100644 --- a/tests/test_litellm/router_strategy/test_least_busy.py +++ b/tests/test_litellm/router_strategy/test_least_busy.py @@ -133,14 +133,51 @@ class UnavailableRedis(SharedRedisCounters): raise ConnectionError("redis is down") -def test_redis_outage_never_fails_the_request() -> None: +@pytest.mark.asyncio +async def test_a_redis_outage_falls_back_to_this_workers_own_counts() -> None: worker: Final = _worker(UnavailableRedis()) worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a")) + assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B + assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B + + await worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None) + assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A +def test_a_shared_counter_that_expired_mid_request_cannot_go_negative() -> None: + shared: Final = SharedRedisCounters() + worker: Final = _worker(shared) + + worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a")) + shared.encoded.clear() + worker.log_success_event(_call_kwargs("dep-a"), None, None, None) + + assert shared.count(f"{GROUP}_request_count:dep-a") == 0 + + worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a")) + + assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B + + +@pytest.mark.asyncio +async def test_a_local_counter_that_expired_mid_request_cannot_go_negative() -> None: + worker: Final = _worker(None) + in_memory: Final = worker.router_cache.in_memory_cache + + worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a")) + in_memory.delete_cache(f"{GROUP}_request_count:dep-a") + await worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None) + + assert worker.router_cache.get_cache(f"{GROUP}_request_count:dep-a") == 0 + + worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a")) + + assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B + + def test_calls_without_a_deployment_are_ignored() -> None: shared: Final = SharedRedisCounters() worker: Final = _worker(shared)