diff --git a/litellm/constants.py b/litellm/constants.py
index b4551a78f5f..eb59858d433 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -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
diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html
deleted file mode 100644
index 82f43619df3..00000000000
--- a/litellm/proxy/_experimental/out/onboarding.html
+++ /dev/null
@@ -1 +0,0 @@
-
LiteLLM Dashboard
\ No newline at end of file
diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml
index 64100277a80..36ec0550481 100644
--- a/litellm/proxy/_new_secret_config.yaml
+++ b/litellm/proxy/_new_secret_config.yaml
@@ -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
diff --git a/litellm/router_strategy/base_routing_strategy.py b/litellm/router_strategy/base_routing_strategy.py
new file mode 100644
index 00000000000..a39d17e3862
--- /dev/null
+++ b/litellm/router_strategy/base_routing_strategy.py
@@ -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()
diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py
index 667246ea2f4..d1a46b7ea89 100644
--- a/litellm/router_strategy/lowest_tpm_rpm_v2.py
+++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py
@@ -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(
diff --git a/tests/litellm/router_strategy/test_base_routing_strategy.py b/tests/litellm/router_strategy/test_base_routing_strategy.py
new file mode 100644
index 00000000000..b47a2f1c90f
--- /dev/null
+++ b/tests/litellm/router_strategy/test_base_routing_strategy.py
@@ -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