mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
fix(router_strategy): stop sharing mutable default dictionaries across handler instances
This commit is contained in:
parent
c2c2a623c0
commit
c598dff3db
5 changed files with 76 additions and 21 deletions
|
|
@ -15,8 +15,9 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
logged_success: int = 0
|
||||
logged_failure: int = 0
|
||||
|
||||
def __init__(self, router_cache: DualCache, routing_args: dict = {}):
|
||||
def __init__(self, router_cache: DualCache, routing_args: dict | None = None):
|
||||
self.router_cache = router_cache
|
||||
self.routing_args = routing_args or {}
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -53,9 +53,9 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
logged_success: int = 0
|
||||
logged_failure: int = 0
|
||||
|
||||
def __init__(self, router_cache: DualCache, routing_args: dict = {}):
|
||||
def __init__(self, router_cache: DualCache, routing_args: dict | None = None):
|
||||
self.router_cache = router_cache
|
||||
self.routing_args = RoutingArgs(**routing_args)
|
||||
self.routing_args = RoutingArgs(**(routing_args or {}))
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -22,9 +22,9 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
logged_failure: int = 0
|
||||
default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour
|
||||
|
||||
def __init__(self, router_cache: DualCache, routing_args: dict = {}):
|
||||
def __init__(self, router_cache: DualCache, routing_args: dict | None = None):
|
||||
self.router_cache = router_cache
|
||||
self.routing_args = RoutingArgs(**routing_args)
|
||||
self.routing_args = RoutingArgs(**(routing_args or {}))
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -48,9 +48,9 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
logged_failure: int = 0
|
||||
default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour
|
||||
|
||||
def __init__(self, router_cache: DualCache, routing_args: dict = {}):
|
||||
def __init__(self, router_cache: DualCache, routing_args: dict | None = None):
|
||||
self.router_cache = router_cache
|
||||
self.routing_args = RoutingArgs(**routing_args)
|
||||
self.routing_args = RoutingArgs(**(routing_args or {}))
|
||||
BaseRoutingStrategy.__init__(
|
||||
self,
|
||||
dual_cache=router_cache,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,10 @@ from unittest.mock import MagicMock, patch
|
|||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError, RedisPipelineIncrementOperation
|
||||
from litellm.router_strategy.base_routing_strategy import BaseRoutingStrategy
|
||||
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
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -60,9 +64,7 @@ async def test_increment_value_in_current_window(base_strategy, mock_dual_cache)
|
|||
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
|
||||
)
|
||||
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
|
||||
|
|
@ -102,9 +104,7 @@ async def test_sync_in_memory_spend_with_redis(base_strategy, mock_dual_cache):
|
|||
# Mock the in-memory cache batch get responses for before snapshot
|
||||
in_memory_before_future: asyncio.Future[List[str]] = asyncio.Future()
|
||||
in_memory_before_future.set_result(["5.0"]) # Initial values
|
||||
mock_dual_cache.in_memory_cache.async_batch_get_cache.return_value = (
|
||||
in_memory_before_future
|
||||
)
|
||||
mock_dual_cache.in_memory_cache.async_batch_get_cache.return_value = in_memory_before_future
|
||||
|
||||
# Mock Redis batch get response
|
||||
redis_future: asyncio.Future[Dict[str, str]] = asyncio.Future()
|
||||
|
|
@ -114,19 +114,14 @@ async def test_sync_in_memory_spend_with_redis(base_strategy, mock_dual_cache):
|
|||
# Mock in-memory get for after snapshot
|
||||
in_memory_after_future: asyncio.Future[Optional[str]] = asyncio.Future()
|
||||
in_memory_after_future.set_result("8.0") # Value after potential updates
|
||||
mock_dual_cache.in_memory_cache.async_get_cache.return_value = (
|
||||
in_memory_after_future
|
||||
)
|
||||
mock_dual_cache.in_memory_cache.async_get_cache.return_value = in_memory_after_future
|
||||
|
||||
await base_strategy._sync_in_memory_spend_with_redis()
|
||||
|
||||
# Verify the final merged values
|
||||
set_cache_calls = mock_dual_cache.in_memory_cache.async_set_cache.call_args_list
|
||||
print(f"set_cache_calls: {set_cache_calls}")
|
||||
assert any(
|
||||
call.kwargs["key"] == "key1" and float(call.kwargs["value"]) == 18.0
|
||||
for call in set_cache_calls
|
||||
)
|
||||
assert any(call.kwargs["key"] == "key1" and float(call.kwargs["value"]) == 18.0 for call in set_cache_calls)
|
||||
|
||||
# Verify cache keys still exist
|
||||
assert len(base_strategy.in_memory_keys_to_update) == 1
|
||||
|
|
@ -150,7 +145,9 @@ async def test_cache_keys_management(base_strategy):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_push_refused_by_the_open_circuit_breaker_is_not_logged_as_an_error(base_strategy, mock_dual_cache, caplog):
|
||||
async def test_push_refused_by_the_open_circuit_breaker_is_not_logged_as_an_error(
|
||||
base_strategy, mock_dual_cache, caplog
|
||||
):
|
||||
"""The sync loop pushes every 100 ms under usage-based routing, so an open breaker must not add an error line per cycle."""
|
||||
mock_dual_cache.redis_cache.async_increment_pipeline.side_effect = RedisCircuitBreakerOpenError(
|
||||
"Redis circuit breaker is open - skipping async_increment_pipeline"
|
||||
|
|
@ -162,3 +159,60 @@ async def test_push_refused_by_the_open_circuit_breaker_is_not_logged_as_an_erro
|
|||
|
||||
assert caplog.records == []
|
||||
assert base_strategy.redis_increment_operation_queue == []
|
||||
|
||||
|
||||
class TestRouterStrategyInstanceIsolation:
|
||||
"""Behavioral regression tests ensuring strategy handlers maintain isolated state when initialized with default arguments."""
|
||||
|
||||
def test_lowest_cost_isolated_routing_args(self) -> None:
|
||||
"""Verify LowestCostLoggingHandler instances do not share routing_args dict."""
|
||||
mock_cache = MagicMock(spec=DualCache)
|
||||
h1 = LowestCostLoggingHandler(router_cache=mock_cache)
|
||||
h2 = LowestCostLoggingHandler(router_cache=mock_cache)
|
||||
|
||||
assert h1.routing_args is not h2.routing_args
|
||||
assert h1.routing_args == {}
|
||||
|
||||
h1.routing_args["ttl"] = 999
|
||||
assert "ttl" not in h2.routing_args
|
||||
|
||||
def test_lowest_latency_isolated_routing_args(self) -> None:
|
||||
"""Verify LowestLatencyLoggingHandler instances initialize independent RoutingArgs objects."""
|
||||
mock_cache = MagicMock(spec=DualCache)
|
||||
h1 = LowestLatencyLoggingHandler(router_cache=mock_cache)
|
||||
h2 = LowestLatencyLoggingHandler(router_cache=mock_cache)
|
||||
|
||||
assert h1.routing_args is not h2.routing_args
|
||||
assert h1.routing_args.ttl == 3600
|
||||
|
||||
h1.routing_args.ttl = 120
|
||||
assert h2.routing_args.ttl == 3600
|
||||
|
||||
def test_lowest_tpm_isolated_routing_args(self) -> None:
|
||||
"""Verify LowestTPMLoggingHandler instances initialize independent RoutingArgs objects."""
|
||||
mock_cache = MagicMock(spec=DualCache)
|
||||
h1 = LowestTPMLoggingHandler(router_cache=mock_cache)
|
||||
h2 = LowestTPMLoggingHandler(router_cache=mock_cache)
|
||||
|
||||
assert h1.routing_args is not h2.routing_args
|
||||
assert h1.routing_args.ttl == 60
|
||||
|
||||
h1.routing_args.ttl = 300
|
||||
assert h2.routing_args.ttl == 60
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lowest_tpm_v2_isolated_routing_args(self) -> None:
|
||||
"""Verify LowestTPMLoggingHandler_v2 instances initialize independent RoutingArgs objects."""
|
||||
mock_cache = MagicMock(spec=DualCache)
|
||||
h1 = LowestTPMLoggingHandler_v2(router_cache=mock_cache)
|
||||
h2 = LowestTPMLoggingHandler_v2(router_cache=mock_cache)
|
||||
|
||||
try:
|
||||
assert h1.routing_args is not h2.routing_args
|
||||
assert h1.routing_args.ttl == 60
|
||||
|
||||
h1.routing_args.ttl = 300
|
||||
assert h2.routing_args.ttl == 60
|
||||
finally:
|
||||
await h1.cleanup()
|
||||
await h2.cleanup()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue