Merge pull request #9357 from BerriAI/litellm_dev_03_18_2025_p2

fix(lowest_tpm_rpm_v2.py): support batch writing increments to redis
This commit is contained in:
Krish Dholakia 2025-03-19 15:45:10 -07:00 • committed by GitHub
commit 9432d1a865
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 341 additions and 9 deletions

View file

@ -7,6 +7,7 @@ DEFAULT_MAX_RETRIES = 2
DEFAULT_FAILURE_THRESHOLD_PERCENT = (
0.5 # default cooldown a deployment if 50% of requests fail in a given minute
)
DEFAULT_REDIS_SYNC_INTERVAL = 1
DEFAULT_COOLDOWN_TIME_SECONDS = 5
DEFAULT_REPLICATE_POLLING_RETRIES = 5
DEFAULT_REPLICATE_POLLING_DELAY_SECONDS = 1

File diff suppressed because one or more lines are too long

View file

@ -3,10 +3,12 @@ model_list:
litellm_params:
model: azure/chatgpt-v-2
api_key: os.environ/AZURE_API_KEY
api_base: os.environ/AZURE_API_BASE
api_base: http://0.0.0.0:8090
rpm: 3
litellm_settings:
callbacks: ["prometheus"]
num_retries: 0
callbacks: ["otel"]
router_settings:
routing_strategy: usage-based-routing-v2 # 👈 KEY CHANGE

View file

@ -0,0 +1,190 @@
"""
Base class across routing strategies to abstract commmon functions like batch incrementing redis
"""
import asyncio
import threading
from abc import ABC
from typing import List, Optional, Set, Union
from litellm._logging import verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import RedisPipelineIncrementOperation
from litellm.constants import DEFAULT_REDIS_SYNC_INTERVAL
class BaseRoutingStrategy(ABC):
def __init__(
self,
dual_cache: DualCache,
should_batch_redis_writes: bool,
default_sync_interval: Optional[Union[int, float]],
):
self.dual_cache = dual_cache
self.redis_increment_operation_queue: List[RedisPipelineIncrementOperation] = []
if should_batch_redis_writes:
try:
# Try to get existing event loop
loop = asyncio.get_event_loop()
if loop.is_running():
# If loop exists and is running, create task in existing loop
loop.create_task(
self.periodic_sync_in_memory_spend_with_redis(
default_sync_interval=default_sync_interval
)
)
else:
self._create_sync_thread(default_sync_interval)
except RuntimeError: # No event loop in current thread
self._create_sync_thread(default_sync_interval)
self.in_memory_keys_to_update: set[str] = (
set()
) # Set with max size of 1000 keys
async def _increment_value_in_current_window(
self, key: str, value: Union[int, float], ttl: int
):
"""
Increment spend within existing budget window
Runs once the budget start time exists in Redis Cache (on the 2nd and subsequent requests to the same provider)
- Increments the spend in memory cache (so spend instantly updated in memory)
- Queues the increment operation to Redis Pipeline (using batched pipeline to optimize performance. Using Redis for multi instance environment of LiteLLM)
"""
result = await self.dual_cache.in_memory_cache.async_increment(
key=key,
value=value,
ttl=ttl,
)
increment_op = RedisPipelineIncrementOperation(
key=key,
increment_value=value,
ttl=ttl,
)
self.redis_increment_operation_queue.append(increment_op)
self.add_to_in_memory_keys_to_update(key=key)
return result
async def periodic_sync_in_memory_spend_with_redis(
self, default_sync_interval: Optional[Union[int, float]]
):
"""
Handler that triggers sync_in_memory_spend_with_redis every DEFAULT_REDIS_SYNC_INTERVAL seconds
Required for multi-instance environment usage of provider budgets
"""
default_sync_interval = default_sync_interval or DEFAULT_REDIS_SYNC_INTERVAL
while True:
try:
await self._sync_in_memory_spend_with_redis()
await asyncio.sleep(
default_sync_interval
) # Wait for DEFAULT_REDIS_SYNC_INTERVAL seconds before next sync
except Exception as e:
verbose_router_logger.error(f"Error in periodic sync task: {str(e)}")
await asyncio.sleep(
default_sync_interval
) # Still wait DEFAULT_REDIS_SYNC_INTERVAL seconds on error before retrying
async def _push_in_memory_increments_to_redis(self):
"""
How this works:
- async_log_success_event collects all provider spend increments in `redis_increment_operation_queue`
- This function pushes all increments to Redis in a batched pipeline to optimize performance
Only runs if Redis is initialized
"""
try:
if not self.dual_cache.redis_cache:
return # Redis is not initialized
verbose_router_logger.debug(
"Pushing Redis Increment Pipeline for queue: %s",
self.redis_increment_operation_queue,
)
if len(self.redis_increment_operation_queue) > 0:
asyncio.create_task(
self.dual_cache.redis_cache.async_increment_pipeline(
increment_list=self.redis_increment_operation_queue,
)
)
self.redis_increment_operation_queue = []
except Exception as e:
verbose_router_logger.error(
f"Error syncing in-memory cache with Redis: {str(e)}"
)
self.redis_increment_operation_queue = []
def add_to_in_memory_keys_to_update(self, key: str):
self.in_memory_keys_to_update.add(key)
def get_in_memory_keys_to_update(self) -> Set[str]:
return self.in_memory_keys_to_update
def reset_in_memory_keys_to_update(self):
self.in_memory_keys_to_update = set()
async def _sync_in_memory_spend_with_redis(self):
"""
Ensures in-memory cache is updated with latest Redis values for all provider spends.
Why Do we need this?
- Optimization to hit sub 100ms latency. Performance was impacted when redis was used for read/write per request
- Use provider budgets in multi-instance environment, we use Redis to sync spend across all instances
What this does:
1. Push all provider spend increments to Redis
2. Fetch all current provider spend from Redis to update in-memory cache
"""
try:
# No need to sync if Redis cache is not initialized
if self.dual_cache.redis_cache is None:
return
# 1. Push all provider spend increments to Redis
await self._push_in_memory_increments_to_redis()
# 2. Fetch all current provider spend from Redis to update in-memory cache
cache_keys = self.get_in_memory_keys_to_update()
cache_keys_list = list(cache_keys)
# Batch fetch current spend values from Redis
redis_values = await self.dual_cache.redis_cache.async_batch_get_cache(
key_list=cache_keys_list
)
# Update in-memory cache with Redis values
if isinstance(redis_values, dict): # Check if redis_values is a dictionary
for key, value in redis_values.items():
if value is not None:
await self.dual_cache.in_memory_cache.async_set_cache(
key=key, value=float(value)
)
verbose_router_logger.debug(
f"Updated in-memory cache for {key}: {value}"
)
self.reset_in_memory_keys_to_update()
except Exception as e:
verbose_router_logger.exception(
f"Error syncing in-memory cache with Redis: {str(e)}"
)
def _create_sync_thread(self, default_sync_interval):
"""Helper method to create a new thread for periodic sync"""
thread = threading.Thread(
target=asyncio.run,
args=(
self.periodic_sync_in_memory_spend_with_redis(
default_sync_interval=default_sync_interval
),
),
daemon=True,
)
thread.start()

View file

@ -15,6 +15,8 @@ from litellm.types.router import RouterErrors
from litellm.types.utils import LiteLLMPydanticObjectBase, StandardLoggingPayload
from litellm.utils import get_utc_datetime, print_verbose
from .base_routing_strategy import BaseRoutingStrategy
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -27,7 +29,7 @@ class RoutingArgs(LiteLLMPydanticObjectBase):
ttl: int = 1 * 60 # 1min (RPM/TPM expire key)
class LowestTPMLoggingHandler_v2(CustomLogger):
class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
"""
Updated version of TPM/RPM Logging.
@ -51,6 +53,12 @@ class LowestTPMLoggingHandler_v2(CustomLogger):
self.router_cache = router_cache
self.model_list = model_list
self.routing_args = RoutingArgs(**routing_args)
BaseRoutingStrategy.__init__(
self,
dual_cache=router_cache,
should_batch_redis_writes=True,
default_sync_interval=0.1,
)
def pre_call_check(self, deployment: Dict) -> Optional[Dict]:
"""
@ -107,6 +115,7 @@ class LowestTPMLoggingHandler_v2(CustomLogger):
)
else:
# if local result below limit, check redis ## prevent unnecessary redis checks
result = self.router_cache.increment_cache(
key=rpm_key, value=1, ttl=self.routing_args.ttl
)
@ -191,11 +200,8 @@ class LowestTPMLoggingHandler_v2(CustomLogger):
)
else:
# if local result below limit, check redis ## prevent unnecessary redis checks
result = await self.router_cache.async_increment_cache(
key=rpm_key,
value=1,
ttl=self.routing_args.ttl,
parent_otel_span=parent_otel_span,
result = await self._increment_value_in_current_window(
key=rpm_key, value=1, ttl=self.routing_args.ttl
)
if result is not None and result > deployment_rpm:
raise litellm.RateLimitError(

View file

@ -0,0 +1,134 @@
import json
import os
import sys
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import asyncio
from unittest.mock import MagicMock, patch
import pytest
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import RedisPipelineIncrementOperation
from litellm.router_strategy.base_routing_strategy import BaseRoutingStrategy
@pytest.fixture
def mock_dual_cache():
dual_cache = MagicMock(spec=DualCache)
dual_cache.in_memory_cache = MagicMock()
dual_cache.redis_cache = MagicMock()
# Set up async method mocks to return coroutines
future1 = asyncio.Future()
future1.set_result(None)
dual_cache.in_memory_cache.async_increment.return_value = future1
future2 = asyncio.Future()
future2.set_result(None)
dual_cache.redis_cache.async_increment_pipeline.return_value = future2
future3 = asyncio.Future()
future3.set_result(None)
dual_cache.in_memory_cache.async_set_cache.return_value = future3
# Fix for async_batch_get_cache
batch_future = asyncio.Future()
batch_future.set_result({"key1": "10.0", "key2": "20.0"})
dual_cache.redis_cache.async_batch_get_cache.return_value = batch_future
return dual_cache
@pytest.fixture
def base_strategy(mock_dual_cache):
return BaseRoutingStrategy(
dual_cache=mock_dual_cache,
should_batch_redis_writes=False,
default_sync_interval=1,
)
@pytest.mark.asyncio
async def test_increment_value_in_current_window(base_strategy, mock_dual_cache):
# Test incrementing value in current window
key = "test_key"
value = 10.0
ttl = 3600
await base_strategy._increment_value_in_current_window(key, value, ttl)
# Verify in-memory cache was incremented
mock_dual_cache.in_memory_cache.async_increment.assert_called_once_with(
key=key, value=value, ttl=ttl
)
# Verify operation was queued for Redis
assert len(base_strategy.redis_increment_operation_queue) == 1
queued_op = base_strategy.redis_increment_operation_queue[0]
assert isinstance(queued_op, dict)
assert queued_op["key"] == key
assert queued_op["increment_value"] == value
assert queued_op["ttl"] == ttl
@pytest.mark.asyncio
async def test_push_in_memory_increments_to_redis(base_strategy, mock_dual_cache):
# Add some operations to the queue
base_strategy.redis_increment_operation_queue = [
RedisPipelineIncrementOperation(key="key1", increment_value=10, ttl=3600),
RedisPipelineIncrementOperation(key="key2", increment_value=20, ttl=3600),
]
await base_strategy._push_in_memory_increments_to_redis()
# Verify Redis pipeline was called
mock_dual_cache.redis_cache.async_increment_pipeline.assert_called_once()
# Verify queue was cleared
assert len(base_strategy.redis_increment_operation_queue) == 0
@pytest.mark.asyncio
async def test_sync_in_memory_spend_with_redis(base_strategy, mock_dual_cache):
# Setup test data
base_strategy.in_memory_keys_to_update = {"key1", "key2"}
# No need to set return_value here anymore as it's set in the fixture
await base_strategy._sync_in_memory_spend_with_redis()
# Verify Redis batch get was called with sorted list for consistent testing
key_list = mock_dual_cache.redis_cache.async_batch_get_cache.call_args.kwargs[
"key_list"
]
sorted(key_list) == sorted(["key1", "key2"])
# mock_dual_cache.redis_cache.async_batch_get_cache.assert_called_once_with(
# key_list=sorted()
# )
# Verify in-memory cache was updated
assert mock_dual_cache.in_memory_cache.async_set_cache.call_count == 2
# Verify cache keys were reset
assert len(base_strategy.in_memory_keys_to_update) == 0
def test_cache_keys_management(base_strategy):
# Test adding and getting cache keys
base_strategy.add_to_in_memory_keys_to_update("key1")
base_strategy.add_to_in_memory_keys_to_update("key2")
base_strategy.add_to_in_memory_keys_to_update("key1") # Duplicate should be ignored
cache_keys = base_strategy.get_in_memory_keys_to_update()
assert len(cache_keys) == 2
assert "key1" in cache_keys
assert "key2" in cache_keys
# Test resetting cache keys
base_strategy.reset_in_memory_keys_to_update()
assert len(base_strategy.get_in_memory_keys_to_update()) == 0