mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
9432d1a865
6 changed files with 341 additions and 9 deletions
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
190
litellm/router_strategy/base_routing_strategy.py
Normal file
190
litellm/router_strategy/base_routing_strategy.py
Normal 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()
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
134
tests/litellm/router_strategy/test_base_routing_strategy.py
Normal file
134
tests/litellm/router_strategy/test_base_routing_strategy.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue