Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
172 changes: 127 additions & 45 deletions litellm/caching/dual_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,9 +8,11 @@
- 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

Expand Down Expand Up @@ -47,6 +49,16 @@ def __setitem__(self, key, value):
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.
Expand Down Expand Up @@ -301,6 +313,37 @@ def _rollback_redis_batch_key_reservations(self, previous_access_times: dict[str
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,
Expand All @@ -309,51 +352,22 @@ async def async_batch_get_cache(
**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,
Expand All @@ -363,6 +377,74 @@ async def async_batch_get_cache(
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:
Expand Down
10 changes: 7 additions & 3 deletions litellm/integrations/SlackAlerting/slack_alerting.py
Original file line number Diff line number Diff line change
Expand Up @@ -376,17 +376,21 @@ async def send_daily_reports(self, router) -> bool:
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

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)
Expand Down
26 changes: 23 additions & 3 deletions litellm/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,7 +134,7 @@
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,
Expand Down Expand Up @@ -259,6 +259,7 @@
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 (
Expand Down Expand Up @@ -1782,6 +1783,7 @@ async def _select_deployment_async(
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
Expand All @@ -1807,6 +1809,14 @@ async def _select_deployment_async(
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,
Expand Down Expand Up @@ -12864,6 +12874,7 @@ async def async_get_healthy_deployments(
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.
Expand Down Expand Up @@ -12916,8 +12927,14 @@ async def async_get_healthy_deployments(
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)
Expand Down Expand Up @@ -13195,6 +13212,7 @@ async def async_get_available_deployment(
# 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,
Expand All @@ -13203,6 +13221,7 @@ async def async_get_available_deployment(
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(
Expand Down Expand Up @@ -13233,6 +13252,7 @@ async def async_get_available_deployment(
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(
Expand Down
Loading
Loading