fix(router_strategy): stop sharing mutable default dictionaries across handler instances

This commit is contained in:
Rohit Kanithi 2026-09-13 02:20:14 -05:00
parent c2c2a623c0
commit c598dff3db
5 changed files with 76 additions and 21 deletions

View file

@ -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:

View file

@ -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:

View file

@ -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:

View file

@ -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,

View file

@ -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()