diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 66be77dbb40..3e13848db02 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -8,9 +8,11 @@ Has 4 primary methods: - async_get_cache """ +import itertools import logging import time from collections.abc import Sequence +from dataclasses import dataclass from threading import Lock from typing import TYPE_CHECKING, Any, Final @@ -47,6 +49,16 @@ class LimitedSizeOrderedDict(OrderedDict): super().__setitem__(key, value) +@dataclass(frozen=True) +class PendingBatchRead: + """A batch read that has consulted the in-memory tier and reserved its Redis keys, but not hit Redis yet.""" + + keys: list[str] + result: list[object | None] + redis_keys: list[str] + previous_access_times: dict[str, float | None] + + class DualCache(BaseCache): """ DualCache is a cache implementation that updates both Redis and an in-memory cache simultaneously. @@ -301,6 +313,37 @@ class DualCache(BaseCache): else: self.last_redis_batch_access_time[key] = previous_time + async def _prepare_batch_get(self, keys: list[str], local_only: bool, **kwargs: object) -> PendingBatchRead: + result: list[object | None] = [None] * len(keys) + if self.in_memory_cache is not None: + in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs) + + if in_memory_result is not None: + result = in_memory_result + + redis_keys: list[str] = [] + previous_access_times: dict[str, float | None] = {} + if None in result and self.redis_cache is not None and local_only is False: + redis_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result) + return PendingBatchRead( + keys=keys, result=result, redis_keys=redis_keys, previous_access_times=previous_access_times + ) + + async def _apply_batch_get( + self, pending: PendingBatchRead, redis_result: dict[str, object] | None, **kwargs: object + ) -> list[object | None]: + if redis_result is None or all(v is None for v in redis_result.values()): + return pending.result + + merged: Final[list[object | None]] = [ + redis_result.get(key, value) for key, value in zip(pending.keys, pending.result) + ] + if self.in_memory_cache is not None: + for key, value in redis_result.items(): + if value is not None: + await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs)) + return merged + async def async_batch_get_cache( self, keys: list, @@ -309,51 +352,22 @@ class DualCache(BaseCache): **kwargs, ): try: - result = [None] * len(keys) - if self.in_memory_cache is not None: - in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs) - - if in_memory_result is not None: - result = in_memory_result - - if None in result and self.redis_cache is not None and local_only is False: - """ - - for the none values in the result - - check the redis cache - """ - current_time: Final = time.time() - sublist_keys, previous_access_times = self._reserve_redis_batch_keys(current_time, keys, result) - - # Only hit Redis if enough time has passed since last access. - if len(sublist_keys) > 0: - try: - # If not found in in-memory cache, try fetching from Redis - redis_result: Final = await self.redis_cache.async_batch_get_cache( - sublist_keys, parent_otel_span=parent_otel_span - ) - except Exception as e: - # Do not throttle subsequent callers if the Redis read fails. - self._rollback_redis_batch_key_reservations(previous_access_times) - if isinstance(e, RedisCircuitBreakerOpenError): - verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e) - return result - raise - - # Short-circuit if redis_result is None or contains only None values - if redis_result is None or all(v is None for v in redis_result.values()): - return result - - # Pre-compute key-to-index mapping for O(1) lookup - key_to_index: Final = {key: i for i, key in enumerate(keys)} - - # Update both result and in-memory cache in a single loop - for key, value in redis_result.items(): - result[key_to_index[key]] = value - - if value is not None and self.in_memory_cache is not None: - await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs)) - - return result + pending: Final = await self._prepare_batch_get(keys, local_only, **kwargs) + # Only hit Redis for keys memory could not serve and enough time has passed since last access. + if not pending.redis_keys or self.redis_cache is None: + return pending.result + try: + redis_result: Final = await self.redis_cache.async_batch_get_cache( + pending.redis_keys, parent_otel_span=parent_otel_span + ) + except Exception as e: + # Do not throttle subsequent callers if the Redis read fails. + self._rollback_redis_batch_key_reservations(pending.previous_access_times) + if isinstance(e, RedisCircuitBreakerOpenError): + verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e) + return pending.result + raise + return await self._apply_batch_get(pending, redis_result, **kwargs) except Exception as e: log_redis_failure( verbose_logger, @@ -363,6 +377,74 @@ class DualCache(BaseCache): with_traceback=True, ) + @staticmethod + async def async_batch_get_cache_shared( + reads: Sequence[tuple["DualCache", list[str]]], + parent_otel_span: Span | None = None, + ) -> list[list[object | None] | None]: + """ + `async_batch_get_cache` for several caches in one Redis round trip. + + Each cache still serves what it can from its own in-memory tier, applies its own Redis read + throttle and backfills its own memory; only the Redis MGET is shared. A failed MGET is reported + to every cache that took part in it exactly as its own failed `async_batch_get_cache` would be: + None when the read raised, the in-memory result when the circuit breaker is open. A cache whose + Redis client is not the one the first cache uses falls back to its own read. + """ + results: Final[list[list[object | None] | None]] = [None] * len(reads) + shared_redis: Final = reads[0][0].redis_cache if reads else None + pendings: Final[list[tuple[int, DualCache, PendingBatchRead]]] = [] + for index, (cache, keys) in enumerate(reads): + if shared_redis is None or cache.redis_cache is not shared_redis: + results[index] = await cache.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) + continue + try: + pending = await cache._prepare_batch_get(keys, local_only=False) + except Exception as e: + DualCache._log_shared_batch_get_failure(e) + continue + pendings.append((index, cache, pending)) + results[index] = pending.result + + redis_keys: Final = list( + dict.fromkeys(itertools.chain.from_iterable(pending.redis_keys for _, _, pending in pendings)) + ) + if shared_redis is None or not redis_keys: + return results + try: + redis_result: Final = await shared_redis.async_batch_get_cache( + redis_keys, parent_otel_span=parent_otel_span + ) + except Exception as e: + for index, cache, pending in pendings: + cache._rollback_redis_batch_key_reservations(pending.previous_access_times) + if pending.redis_keys and not isinstance(e, RedisCircuitBreakerOpenError): + results[index] = None + if isinstance(e, RedisCircuitBreakerOpenError): + verbose_logger.debug("LiteLLM Cache: async_batch_get_cache_shared served from memory only: %s", e) + else: + DualCache._log_shared_batch_get_failure(e) + return results + + for index, cache, pending in pendings: + own_result = {key: redis_result[key] for key in pending.redis_keys if key in redis_result} + try: + results[index] = await cache._apply_batch_get(pending, own_result) + except Exception as e: + results[index] = None + DualCache._log_shared_batch_get_failure(e) + return results + + @staticmethod + def _log_shared_batch_get_failure(e: Exception) -> None: + log_redis_failure( + verbose_logger, + logging.ERROR, + "LiteLLM Cache: exception in async_batch_get_cache_shared", + e, + with_traceback=True, + ) + async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): print_verbose(f"async set cache: cache key: {key}; local_only: {local_only}; value: {value}") try: diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 7c608aac8d9..50f63316a62 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -376,8 +376,12 @@ class SlackAlerting(CustomBatchLogger): if combined_metrics_values is None: return False + metric_values: Final[list[float | None]] = [ + val if isinstance(val, (int, float)) else None for val in combined_metrics_values + ] + all_none = True - for val in combined_metrics_values: + for val in metric_values: if val is not None and val > 0: all_none = False break @@ -385,8 +389,8 @@ class SlackAlerting(CustomBatchLogger): if all_none: return False - failed_request_values: Final = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..] - latency_values: Final = combined_metrics_values[len(failed_request_keys) :] + failed_request_values: Final = metric_values[: len(failed_request_keys)] # # [1, 2, None, ..] + latency_values: Final = metric_values[len(failed_request_keys) :] # find top 5 failed ## Replace None values with a placeholder value (-1 in this case) diff --git a/litellm/router.py b/litellm/router.py index 86a67a8d5ca..a98631b7f97 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -134,7 +134,7 @@ from litellm.router_strategy.least_busy import LeastBusyLoggingHandler from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler -from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2 +from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage from litellm.router_strategy.simple_shuffle import simple_shuffle from litellm.router_strategy.tag_based_routing import ( _get_tags_from_request_kwargs, @@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import ( parse_routing_groups, validate_routing_strategy, ) +from litellm.router_utils.routing_read_batch import RoutingReadBatch from litellm.scheduler import FlowItem, Scheduler from litellm.types.litellm_params import RoutingStrategyName from litellm.types.llms.openai import ( @@ -1789,6 +1790,7 @@ class Router: messages: list[dict[str, str]] | None, input: str | list | None, request_kwargs: dict | None, + prefetched_usage: PrefetchedUsage | None = None, ) -> Any | None: """ Asks the strategy selector for a deployment. Caller handles @@ -1814,6 +1816,14 @@ class Router: messages=messages, input=input, ) + case "usage-based-routing-v2" if isinstance(selector, LowestTPMLoggingHandler_v2): + return await selector.async_get_available_deployments( + model_group=model, + healthy_deployments=healthy_deployments, + messages=messages, + input=input, + prefetched_usage=prefetched_usage, + ) case "usage-based-routing-v2" | "cost-based-routing": return await selector.async_get_available_deployments( model_group=model, @@ -12925,6 +12935,7 @@ class Router: specific_deployment: bool | None = False, parent_otel_span: Span | None = None, health_check_probe: bool = False, + routing_read_batch: RoutingReadBatch | None = None, ) -> list[dict] | dict: """ Get the healthy deployments for a model. @@ -12977,8 +12988,14 @@ class Router: health_check_probe=health_check_probe, ) - cooldown_deployments: Final = await _async_get_cooldown_deployments( - litellm_router_instance=self, parent_otel_span=parent_otel_span + cooldown_deployments: Final = ( + await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) + if routing_read_batch is None + else await routing_read_batch.async_get_cooldown_deployments( + litellm_router_instance=self, + healthy_deployments=healthy_deployments, + parent_otel_span=parent_otel_span, + ) ) if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("cooldown deployments: %s", cooldown_deployments) @@ -13256,6 +13273,7 @@ class Router: # the hook can replace `model` and routing-group lookup must key # off the final model name. strategy, strategy_selector = self._get_routing_context(model, request_kwargs) + routing_read_batch: Final = RoutingReadBatch.for_strategy(strategy, strategy_selector) healthy_deployments: Final = await self.async_get_healthy_deployments( model=model, @@ -13264,6 +13282,7 @@ class Router: input=input, specific_deployment=specific_deployment, parent_otel_span=parent_otel_span, + routing_read_batch=routing_read_batch, ) if isinstance(healthy_deployments, dict): await self._async_override_selector_pre_call_check( @@ -13294,6 +13313,7 @@ class Router: messages=messages, input=input, request_kwargs=request_kwargs, + prefetched_usage=routing_read_batch.prefetched_usage if routing_read_batch is not None else None, ) if deployment is None: exception: Final = await async_raise_no_deployment_exception( diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index a2acce5fcb5..909b47833cb 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -1,7 +1,8 @@ #### What this does #### # identifies lowest tpm deployment import random -from collections.abc import Sequence +from collections.abc import Mapping, Sequence +from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Final import httpx @@ -31,6 +32,26 @@ class RoutingArgs(LiteLLMPydanticObjectBase): ttl: int = 1 * 60 # 1min (RPM/TPM expire key) +@dataclass(frozen=True) +class PrefetchedUsage: + """ + tpm/rpm counter values another read of this request already fetched from the router cache. + + `values` is None when that read failed, which is what `async_batch_get_cache` returns on failure. + """ + + keys: frozenset[str] + values: Mapping[str, object] | None + + def covers(self, keys: Sequence[str]) -> bool: + return self.keys.issuperset(keys) + + def values_for(self, keys: Sequence[str]) -> list[object | None] | None: + if self.values is None: + return None + return [self.values.get(key) for key in keys] + + class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): """ Updated version of TPM/RPM Logging. @@ -412,17 +433,35 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): else: return None + def usage_counter_keys(self, healthy_deployments: list) -> tuple[list[str], list[str]]: + """The `::tpm:` and `::rpm:` counter keys selection reads.""" + current_minute: Final = get_utc_datetime().strftime("%H-%M") + + tpm_keys: Final[list[str]] = [] + rpm_keys: Final[list[str]] = [] + for m in healthy_deployments: + if isinstance(m, dict): + id = m.get("model_info", {}).get( + "id" + ) # a deployment should always have an 'id'. this is set in router.py + deployment_name = m.get("litellm_params", {}).get("model") + tpm_keys.append(f"{id}:{deployment_name}:tpm:{current_minute}") + rpm_keys.append(f"{id}:{deployment_name}:rpm:{current_minute}") + return tpm_keys, rpm_keys + async def async_get_available_deployments( self, model_group: str, healthy_deployments: list, messages: list[dict[str, str]] | None = None, input: str | list | None = None, + prefetched_usage: PrefetchedUsage | None = None, ): """ Async implementation of get deployments. - Reduces time to retrieve the tpm/rpm values from cache + Reduces time to retrieve the tpm/rpm values from cache. `prefetched_usage` skips the cache + read when it already holds this request's counters (see `RoutingReadBatch`). """ # get list of potential deployments verbose_router_logger.debug( @@ -431,28 +470,15 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): healthy_deployments, ) - dt: Final = get_utc_datetime() - current_minute: Final = dt.strftime("%H-%M") - - tpm_keys: Final = [] - rpm_keys: Final = [] - for m in healthy_deployments: - if isinstance(m, dict): - id = m.get("model_info", {}).get( - "id" - ) # a deployment should always have an 'id'. this is set in router.py - deployment_name = m.get("litellm_params", {}).get("model") - tpm_key = f"{id}:{deployment_name}:tpm:{current_minute}" - rpm_key = f"{id}:{deployment_name}:rpm:{current_minute}" - - tpm_keys.append(tpm_key) - rpm_keys.append(rpm_key) - + tpm_keys, rpm_keys = self.usage_counter_keys(healthy_deployments) combined_tpm_rpm_keys: Final = tpm_keys + rpm_keys - combined_tpm_rpm_values: Final = await self.router_cache.async_batch_get_cache( - keys=combined_tpm_rpm_keys - ) # [1, 2, None, ..] + if prefetched_usage is not None and prefetched_usage.covers(combined_tpm_rpm_keys): + combined_tpm_rpm_values = prefetched_usage.values_for(combined_tpm_rpm_keys) + else: + combined_tpm_rpm_values = await self.router_cache.async_batch_get_cache( + keys=combined_tpm_rpm_keys + ) # [1, 2, None, ..] if combined_tpm_rpm_values is not None: tpm_values = combined_tpm_rpm_values[: len(tpm_keys)] diff --git a/litellm/router_utils/cooldown_cache.py b/litellm/router_utils/cooldown_cache.py index ef29f7d8fd3..187215d3d16 100644 --- a/litellm/router_utils/cooldown_cache.py +++ b/litellm/router_utils/cooldown_cache.py @@ -4,7 +4,7 @@ Wrapper around router cache. Meant to handle model cooldown logic import functools import time -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final from typing_extensions import TypedDict @@ -163,6 +163,12 @@ class CooldownCache: keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids] results: Final = await self.cooldown_store.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) + return self.active_cooldowns_from_results(model_ids, results) + + def active_cooldowns_from_results( + self, model_ids: list[str], results: Sequence[object] | None + ) -> list[tuple[str, CooldownCacheValue]]: + """The cooldowns still active in a `cooldown_store` batch read of `get_cooldown_cache_key(model_id)` per id.""" active_cooldowns: Final[list[tuple[str, CooldownCacheValue]]] = [] if results is None or all(v is None for v in results): diff --git a/litellm/router_utils/routing_read_batch.py b/litellm/router_utils/routing_read_batch.py new file mode 100644 index 00000000000..e9410d0e586 --- /dev/null +++ b/litellm/router_utils/routing_read_batch.py @@ -0,0 +1,72 @@ +""" +One Redis round trip for the reads a request needs before a deployment can be picked. + +The cooldown filter (`CooldownCache`, its own `DualCache`) and usage-based selection +(`LowestTPMLoggingHandler_v2`, the router cache) each issue their own MGET because they live in +different objects. `RoutingReadBatch` fetches both key sets in one +`DualCache.async_batch_get_cache_shared` while the healthy deployments are being resolved and hands +the usage slice to the strategy, so selection does not read again. +""" + +from typing import TYPE_CHECKING, Any, Final + +from litellm._logging import verbose_router_logger +from litellm.caching.dual_cache import DualCache +from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage +from litellm.router_utils.cooldown_cache import CooldownCache + +if TYPE_CHECKING: + from opentelemetry.trace import Span as _Span + + from litellm.router import Router as _Router + + LitellmRouter = _Router + Span = _Span +else: + LitellmRouter = Any + Span = Any + + +class RoutingReadBatch: + def __init__(self, usage_selector: LowestTPMLoggingHandler_v2) -> None: + self.usage_selector: Final = usage_selector + self.prefetched_usage: PrefetchedUsage | None = None + + @staticmethod + def for_strategy(strategy: str | None, selector: object) -> "RoutingReadBatch | None": + if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2): + return RoutingReadBatch(usage_selector=selector) + return None + + async def async_get_cooldown_deployments( + self, + litellm_router_instance: LitellmRouter, + healthy_deployments: list, + parent_otel_span: Span | None, + ) -> list[str]: + """ + `_async_get_cooldown_deployments`, with the strategy's tpm/rpm counters for + `healthy_deployments` fetched in the same MGET and kept as `prefetched_usage`. + """ + model_ids: Final = litellm_router_instance.get_model_ids() + cooldown_keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids] + tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments) + usage_keys: Final = tpm_keys + rpm_keys + + cooldown_results, usage_values = await DualCache.async_batch_get_cache_shared( + [ + (litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys), + (self.usage_selector.router_cache, usage_keys), + ], + parent_otel_span=parent_otel_span, + ) + self.prefetched_usage = PrefetchedUsage( + keys=frozenset(usage_keys), + values=None if usage_values is None else dict(zip(usage_keys, usage_values)), + ) + + cooldown_models: Final = litellm_router_instance.cooldown_cache.active_cooldowns_from_results( + model_ids, cooldown_results + ) + verbose_router_logger.debug("retrieve cooldown models: %s", cooldown_models) + return [model_id for model_id, _ in cooldown_models] diff --git a/tests/unit/caching/test_dual_cache.py b/tests/unit/caching/test_dual_cache.py index 5f59de9cca5..eb2f19ac377 100644 --- a/tests/unit/caching/test_dual_cache.py +++ b/tests/unit/caching/test_dual_cache.py @@ -6,18 +6,16 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache, _redis_circuit_breaker_guard, _redis_circuit_breaker_guard_sync +from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE from litellm.types.caching import RedisPipelineIncrementOperation @pytest.mark.asyncio async def test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads(): - dual_cache = DualCache( - redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10) keys = ["shared_a", "shared_b"] start_gate = asyncio.Event() @@ -44,9 +42,7 @@ async def test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads @pytest.mark.asyncio async def test_dual_cache_async_batch_get_cache_rolls_back_redis_reservation_on_error(): - dual_cache = DualCache( - redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10) keys = ["shared_a", "shared_b"] with patch.object( @@ -116,9 +112,7 @@ def test_dual_cache_batch_get_cache_only_reads_missing_keys_from_redis(): def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads(): mock_redis = _redis_mock_for_sync_batch({"absent_key": None}) - dual_cache = DualCache( - in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10) first = dual_cache.batch_get_cache(keys=["absent_key"]) second = dual_cache.batch_get_cache(keys=["absent_key"]) @@ -131,9 +125,7 @@ def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads(): def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error(): mock_redis = MagicMock(spec=RedisCache) mock_redis.batch_get_cache.side_effect = RuntimeError("redis unavailable") - dual_cache = DualCache( - in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10) first_result = dual_cache.batch_get_cache(keys=["shared_a"]) second_result = dual_cache.batch_get_cache(keys=["shared_a"]) @@ -146,9 +138,7 @@ def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error(): def test_dual_cache_batch_get_cache_returns_memory_only_when_redis_read_is_throttled(): mock_redis = _redis_mock_for_sync_batch({"throttled_key": "redis_value"}) - dual_cache = DualCache( - in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10) dual_cache.last_redis_batch_access_time["throttled_key"] = time.time() result = dual_cache.batch_get_cache(keys=["throttled_key"]) @@ -257,9 +247,7 @@ async def test_dual_cache_batch_redis_backfill_injects_default_in_memory_ttl(): default_in_memory_ttl, same as the single-key path.""" in_memory_cache = InMemoryCache(default_ttl=600) mock_redis = MagicMock(spec=RedisCache) - mock_redis.async_batch_get_cache = AsyncMock( - return_value={"batch_backfill_key": "redis_value"} - ) + mock_redis.async_batch_get_cache = AsyncMock(return_value={"batch_backfill_key": "redis_value"}) dual_cache = DualCache( in_memory_cache=in_memory_cache, redis_cache=mock_redis, @@ -371,9 +359,7 @@ async def test_circuit_breaker_open_skips_redis(): class FakeRedis: def __init__(self): - self._circuit_breaker = RedisCircuitBreaker( - failure_threshold=3, recovery_timeout=60 - ) + self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60) self._circuit_breaker._state = "open" self._circuit_breaker._opened_at = time.time() self.call_count = 0 @@ -426,9 +412,7 @@ def test_circuit_breaker_half_open_concurrent_calls_are_fast_failed(): # All subsequent concurrent callers: HALF_OPEN → fast-fail (return True) for _ in range(10): - assert ( - cb.is_open() is True - ), "concurrent callers should be fast-failed in HALF_OPEN" + assert cb.is_open() is True, "concurrent callers should be fast-failed in HALF_OPEN" def test_circuit_breaker_disabled_never_opens(): @@ -472,9 +456,7 @@ async def test_circuit_breaker_disabled_guard_always_calls_method(): class FakeRedis: def __init__(self): - self._circuit_breaker = RedisCircuitBreaker( - failure_threshold=1, recovery_timeout=60, enabled=False - ) + self._circuit_breaker = RedisCircuitBreaker(failure_threshold=1, recovery_timeout=60, enabled=False) self.call_count = 0 @_redis_circuit_breaker_guard @@ -791,3 +773,125 @@ async def test_async_delete_cache_keys_on_empty_list_touches_no_backend(): await dual_cache.async_delete_cache_keys([]) redis_cache.delete_cache_keys.assert_not_awaited() + + +def _recording_redis(values: dict) -> MagicMock: + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock( + side_effect=lambda key_list, parent_otel_span=None: {key: values.get(key) for key in key_list} + ) + return redis + + +@pytest.mark.asyncio +async def test_shared_batch_read_issues_one_mget_for_two_caches_and_backfills_each_one_separately(): + redis = _recording_redis({"a1": 1, "b2": "x"}) + first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + + results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1", "b2"])]) + + assert results == [[1, None], [None, "x"]] + assert redis.async_batch_get_cache.await_count == 1 + assert redis.async_batch_get_cache.await_args.args[0] == ["a1", "a2", "b1", "b2"] + assert first.in_memory_cache.get_cache("a1") == 1 + assert second.in_memory_cache.get_cache("b2") == "x" + assert first.in_memory_cache.get_cache("b2") is None, "backfill leaked into the other cache" + + +@pytest.mark.asyncio +async def test_shared_batch_read_serves_memory_hits_and_throttles_like_the_separate_reads(): + redis = _recording_redis({"a2": 2, "b1": 3}) + first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + first.in_memory_cache.set_cache("a1", 5) + second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + second.in_memory_cache.set_cache("b1", 3) + + results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1"])]) + + assert results == [[5, 2], [3]] + assert redis.async_batch_get_cache.await_args.args[0] == ["a2"], "memory hits must not hit Redis" + + first.in_memory_cache.delete_cache("a2") + results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1"])]) + + assert results == [[5, None], [3]] + assert redis.async_batch_get_cache.await_count == 1, "a2 was read within the batch expiry, so it is throttled" + + +@pytest.mark.asyncio +async def test_shared_batch_read_failure_degrades_exactly_like_two_failed_reads(): + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable")) + first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + third = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + third.in_memory_cache.set_cache("c1", "memory") + + shared = await DualCache.async_batch_get_cache_shared([(first, ["a1"]), (second, ["b1"]), (third, ["c1"])]) + separate = [ + await first.async_batch_get_cache(keys=["a1"]), + await second.async_batch_get_cache(keys=["b1"]), + await third.async_batch_get_cache(keys=["c1"]), + ] + + assert shared == separate == [None, None, ["memory"]] + assert "a1" not in first.last_redis_batch_access_time + assert "b1" not in second.last_redis_batch_access_time + + +@pytest.mark.asyncio +async def test_shared_batch_read_with_an_open_breaker_keeps_memory_hits_and_releases_reservations(): + first = _dual_cache_with_open_breaker_and_a_memory_hit() + second = DualCache( + in_memory_cache=InMemoryCache(), redis_cache=first.redis_cache, default_redis_batch_cache_expiry=10 + ) + + results = await DualCache.async_batch_get_cache_shared([(first, ["k1", "k2"]), (second, ["k3"])]) + + assert results == [["v1", None], [None]] + assert "k2" not in first.last_redis_batch_access_time + assert "k3" not in second.last_redis_batch_access_time + + +@pytest.mark.asyncio +async def test_shared_batch_read_falls_back_to_a_caches_own_read_when_its_redis_client_differs(): + first_redis = _recording_redis({"a1": 1}) + second_redis = _recording_redis({"b1": 2}) + first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=first_redis, default_redis_batch_cache_expiry=10) + second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=second_redis, default_redis_batch_cache_expiry=10) + memory_only = DualCache(in_memory_cache=InMemoryCache(), redis_cache=None) + memory_only.in_memory_cache.set_cache("m1", "m") + + results = await DualCache.async_batch_get_cache_shared( + [(first, ["a1"]), (second, ["b1"]), (memory_only, ["m1", "m2"])] + ) + + assert results == [[1], [2], ["m", None]] + assert first_redis.async_batch_get_cache.await_args.args[0] == ["a1"] + assert second_redis.async_batch_get_cache.await_args.args[0] == ["b1"] + + +@pytest.mark.asyncio +async def test_shared_batch_read_keeps_a_caches_own_tier_failure_to_itself_like_the_separate_read(): + redis = _recording_redis({"a1": 1, "b1": 2, "c1": 3}) + broken_memory_read = DualCache( + in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10 + ) + broken_memory_read.in_memory_cache.async_batch_get_cache = AsyncMock(side_effect=RuntimeError("memory read")) + broken_backfill = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + broken_backfill.in_memory_cache.async_set_cache = AsyncMock(side_effect=RuntimeError("memory write")) + healthy = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + + shared = await DualCache.async_batch_get_cache_shared( + [(broken_memory_read, ["a1"]), (broken_backfill, ["b1"]), (healthy, ["c1"])] + ) + broken_backfill.last_redis_batch_access_time.clear() + separate = [ + await broken_memory_read.async_batch_get_cache(keys=["a1"]), + await broken_backfill.async_batch_get_cache(keys=["b1"]), + await healthy.async_batch_get_cache(keys=["c1"]), + ] + + assert shared == separate == [None, None, [3]] + assert redis.async_batch_get_cache.await_args_list[0].args[0] == ["b1", "c1"] diff --git a/tests/unit/router_strategy/test_lowest_tpm_rpm.py b/tests/unit/router_strategy/test_lowest_tpm_rpm.py index 7b13b196d5b..0fa11cda20c 100644 --- a/tests/unit/router_strategy/test_lowest_tpm_rpm.py +++ b/tests/unit/router_strategy/test_lowest_tpm_rpm.py @@ -1,7 +1,12 @@ from datetime import datetime, timedelta from typing import Final +from unittest.mock import AsyncMock + +import pytest from litellm import Router +from litellm.caching.dual_cache import DualCache +from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict MODEL_GROUP: Final = "lowest-tpm-router" @@ -52,3 +57,26 @@ def test_usage_based_routing_v1_selects_the_lowest_recorded_tpm() -> None: ) assert deployment["model_info"]["id"] == LOW_USAGE_DEPLOYMENT_ID + + +@pytest.mark.asyncio +async def test_v2_async_selection_uses_prefetched_counters_only_when_they_cover_its_keys(): + router_cache = DualCache() + router_cache.async_batch_get_cache = AsyncMock(return_value=[100, 10, None, None]) # type: ignore[method-assign] + strategy = LowestTPMLoggingHandler_v2(router_cache=router_cache) + deployments = [ + {"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "a"}}, + {"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}}, + ] + tpm_keys, rpm_keys = strategy.usage_counter_keys(deployments) + keys = tpm_keys + rpm_keys + + covering = PrefetchedUsage(keys=frozenset(keys), values=dict(zip(keys, [10, 100, None, None]))) + chosen = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments, prefetched_usage=covering) + assert chosen["model_info"]["id"] == "a", "the prefetched counters say a is the lowest" + router_cache.async_batch_get_cache.assert_not_awaited() + + stale = PrefetchedUsage(keys=frozenset(keys[:1]), values={keys[0]: 10}) + chosen = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments, prefetched_usage=stale) + assert chosen["model_info"]["id"] == "b", "counters that do not cover this minute's keys are read again" + router_cache.async_batch_get_cache.assert_awaited_once_with(keys=keys) diff --git a/tests/unit/router_utils/test_routing_read_batch.py b/tests/unit/router_utils/test_routing_read_batch.py new file mode 100644 index 00000000000..73be5fd4a3a --- /dev/null +++ b/tests/unit/router_utils/test_routing_read_batch.py @@ -0,0 +1,178 @@ +""" +One Redis round trip per request for the router's pre-call reads. + +Before `RoutingReadBatch`, `async_get_available_deployment` issued one MGET for the cooldown keys +(`CooldownCache`) and a second one for the tpm/rpm counters (`LowestTPMLoggingHandler_v2`). +""" + +import time +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm import Router +from litellm.caching.redis_cache import RedisCache + +_MODEL_GROUP = "claude" +_MESSAGES = [{"role": "user", "content": "ping"}] + + +def _deployment(deployment_id: str) -> dict: + return { + "model_name": _MODEL_GROUP, + "litellm_params": {"model": "anthropic/claude-x", "api_key": "test", "mock_response": "pong"}, + "model_info": {"id": deployment_id}, + } + + +def _redis_answering(values_by_key_prefix: dict[str, object]) -> MagicMock: + """A Redis double that answers each key from its minute-less prefix and records every MGET.""" + + def _mget(key_list, parent_otel_span=None): + return {key: values_by_key_prefix.get(key.rsplit(":", 1)[0], values_by_key_prefix.get(key)) for key in key_list} + + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock(side_effect=_mget) + return redis + + +def _router(redis: MagicMock, routing_strategy: str) -> Router: + router = Router( + model_list=[_deployment("dep-a"), _deployment("dep-b")], + routing_strategy=routing_strategy, + ) + router._update_redis_cache(cache=redis) + return router + + +def _redis_key_families(redis: MagicMock) -> list[list[str]]: + return [ + sorted(key.rsplit(":", 1)[0] if ":tpm:" in key or ":rpm:" in key else key for key in call.args[0]) + for call in redis.async_batch_get_cache.await_args_list + ] + + +def _cooldown(seconds: float) -> dict: + return {"exception_received": "429", "status_code": "429", "timestamp": time.time(), "cooldown_time": seconds} + + +@pytest.mark.asyncio +async def test_usage_based_routing_reads_cooldowns_and_counters_in_one_redis_round_trip(): + redis = _redis_answering({}) + router = _router(redis, "usage-based-routing-v2") + + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert _redis_key_families(redis) == [ + [ + "dep-a:anthropic/claude-x:rpm", + "dep-a:anthropic/claude-x:tpm", + "dep-b:anthropic/claude-x:rpm", + "dep-b:anthropic/claude-x:tpm", + "deployment:dep-a:cooldown", + "deployment:dep-b:cooldown", + ] + ], "cooldown state and usage counters must arrive in one MGET" + + +@pytest.mark.asyncio +async def test_simple_shuffle_still_reads_only_cooldowns(): + redis = _redis_answering({}) + router = _router(redis, "simple-shuffle") + + await router.async_get_available_deployment(model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES) + + assert _redis_key_families(redis) == [["deployment:dep-a:cooldown", "deployment:dep-b:cooldown"]] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("tpm_a", "tpm_b", "expected"), + [(100, 10, "dep-b"), (10, 100, "dep-a"), (None, 10, "dep-a"), (10, None, "dep-b")], +) +async def test_batched_counters_pick_the_deployment_the_strategy_picks_reading_alone(tpm_a, tpm_b, expected): + counters = {"dep-a:anthropic/claude-x:tpm": tpm_a, "dep-b:anthropic/claude-x:tpm": tpm_b} + routed = _router(_redis_answering(counters), "usage-based-routing-v2") + alone = _router(_redis_answering(counters), "usage-based-routing-v2") + + routed_choice = await routed.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + alone_choice = await alone.lowesttpm_logger_v2.async_get_available_deployments( + model_group=_MODEL_GROUP, healthy_deployments=alone.model_list, messages=_MESSAGES + ) + + assert routed_choice["model_info"]["id"] == alone_choice["model_info"]["id"] == expected + + +@pytest.mark.asyncio +async def test_batched_read_still_excludes_a_cooled_down_deployment(): + redis = _redis_answering( + { + "dep-a:anthropic/claude-x:tpm": 100, + "dep-b:anthropic/claude-x:tpm": 10, + "deployment:dep-b:cooldown": _cooldown(seconds=60), + } + ) + router = _router(redis, "usage-based-routing-v2") + + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + + assert deployment["model_info"]["id"] == "dep-a", "dep-b has the lowest tpm but is cooling down" + assert redis.async_batch_get_cache.await_count == 1 + + +@pytest.mark.asyncio +async def test_batched_read_ignores_an_expired_cooldown(): + redis = _redis_answering( + { + "dep-a:anthropic/claude-x:tpm": 100, + "dep-b:anthropic/claude-x:tpm": 10, + "deployment:dep-b:cooldown": _cooldown(seconds=-1), + } + ) + router = _router(redis, "usage-based-routing-v2") + + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + + assert deployment["model_info"]["id"] == "dep-b" + + +@pytest.mark.asyncio +async def test_a_failed_batched_read_degrades_like_the_two_failed_reads_did(): + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable")) + routed = _router(redis, "usage-based-routing-v2") + alone = _router(redis, "usage-based-routing-v2") + + with pytest.raises(litellm.RateLimitError, match="No deployments available") as routed_error: + await routed.async_get_available_deployment(model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES) + with pytest.raises(litellm.RateLimitError, match="No deployments available") as alone_error: + await alone.lowesttpm_logger_v2.async_get_available_deployments( + model_group=_MODEL_GROUP, healthy_deployments=alone.model_list, messages=_MESSAGES + ) + + assert str(routed_error.value) == str(alone_error.value) + assert len(routed.cache.last_redis_batch_access_time) == 0, "a failed read must not throttle the next one" + assert len(routed.cooldown_cache.cooldown_store.last_redis_batch_access_time) == 0 + + +@pytest.mark.asyncio +async def test_a_failed_batched_read_leaves_simple_shuffle_routing(): + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable")) + router = _router(redis, "simple-shuffle") + + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"}