diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index 22c321c65fb..39dfa327777 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -16,8 +16,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): if is_batch_retrieve_call_type(kwargs.get("call_type")): diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 66c8227195d..e0cf68c0667 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -54,9 +54,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): if is_batch_retrieve_call_type(kwargs.get("call_type")): diff --git a/litellm/router_strategy/lowest_tpm_rpm.py b/litellm/router_strategy/lowest_tpm_rpm.py index d4abf1f8f70..bd0080d25e9 100644 --- a/litellm/router_strategy/lowest_tpm_rpm.py +++ b/litellm/router_strategy/lowest_tpm_rpm.py @@ -23,9 +23,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): if is_batch_retrieve_call_type(kwargs.get("call_type")): diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index a2acce5fcb5..b790b152d0f 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -49,9 +49,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, diff --git a/tests/test_litellm/router_strategy/test_base_routing_strategy.py b/tests/test_litellm/router_strategy/test_base_routing_strategy.py index dc75c3d2916..551a26f710c 100644 --- a/tests/test_litellm/router_strategy/test_base_routing_strategy.py +++ b/tests/test_litellm/router_strategy/test_base_routing_strategy.py @@ -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,79 @@ 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() + + @pytest.mark.asyncio + async def test_handlers_accept_explicit_routing_args(self) -> None: + """Verify handlers accept and respect explicitly provided routing_args dict.""" + mock_cache = MagicMock(spec=DualCache) + h_cost = LowestCostLoggingHandler(router_cache=mock_cache, routing_args={"custom": True}) + assert h_cost.routing_args == {"custom": True} + + h_lat = LowestLatencyLoggingHandler(router_cache=mock_cache, routing_args={"ttl": 120}) + assert h_lat.routing_args.ttl == 120 + + h_tpm = LowestTPMLoggingHandler(router_cache=mock_cache, routing_args={"ttl": 180}) + assert h_tpm.routing_args.ttl == 180 + + h_tpm_v2 = LowestTPMLoggingHandler_v2(router_cache=mock_cache, routing_args={"ttl": 240}) + try: + assert h_tpm_v2.routing_args.ttl == 240 + finally: + await h_tpm_v2.cleanup()