mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
perf(router): fetch cooldown state and usage counters in one Redis round trip (#43320)
* perf(router): fetch cooldown state and usage counters in one Redis round trip The cooldown filter (CooldownCache) and usage-based-routing-v2 selection (LowestTPMLoggingHandler_v2) each issued their own MGET on every request because they live in different objects. RoutingReadBatch fetches both key sets through DualCache.async_batch_get_cache_shared while the healthy deployments are resolved and hands the usage slice to the strategy, so selection does not read again. Each cache keeps its own memory tier, throttling, reservation rollback and circuit-breaker handling, and the strategy falls back to its own read when the prefetch does not cover its keys. simple-shuffle keeps reading only cooldowns. aresponses no longer issues a second, blocking response-cache read from the worker thread that runs the sync wrapper. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): keep per-cache tier failures inside the shared batch read Wrap the memory-tier prepare and backfill steps of DualCache.async_batch_get_cache_shared so a failing tier degrades that cache's read to None the way async_batch_get_cache does, instead of escaping into routing. Drop the aresponses sync-cache guard: for native Responses models the worker-thread read is the one whose key matches the write, so skipping it broke cached /v1/responses replays. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(router): rename usage key builder so the async cache-call check reads it as a key helper Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(alerting): narrow daily-report cache values before numeric comparison Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): type the shared batch-read helpers and merge Redis results without mutation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style: fix import sort in test_dual_cache Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): flatten shared batch read keys without a stacked comprehension Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
27c110cb71
commit
d2a574b791
9 changed files with 622 additions and 102 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 `<id>:<model>:tpm:<HH-MM>` and `<id>:<model>:rpm:<HH-MM>` 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)]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
72
litellm/router_utils/routing_read_batch.py
Normal file
72
litellm/router_utils/routing_read_batch.py
Normal file
|
|
@ -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]
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
178
tests/unit/router_utils/test_routing_read_batch.py
Normal file
178
tests/unit/router_utils/test_routing_read_batch.py
Normal file
|
|
@ -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"}
|
||||
Loading…
Add table
Reference in a new issue