From 13badfa9ae1d52a54bf5895962cd8acc25450f92 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 31 May 2025 00:45:19 -0700 Subject: [PATCH] test: update testing --- .../router_strategy/base_routing_strategy.py | 2 - .../test_base_routing_strategy.py | 52 ++++++++++++++----- 2 files changed, 39 insertions(+), 15 deletions(-) diff --git a/litellm/router_strategy/base_routing_strategy.py b/litellm/router_strategy/base_routing_strategy.py index 997288813d9..c40149917bb 100644 --- a/litellm/router_strategy/base_routing_strategy.py +++ b/litellm/router_strategy/base_routing_strategy.py @@ -178,8 +178,6 @@ class BaseRoutingStrategy(ABC): await self._push_in_memory_increments_to_redis() # 2. Fetch all current provider spend from Redis to update in-memory cache - pattern = self.get_key_pattern_to_sync() - cache_keys: Optional[Union[Set[str], List[str]]] = None cache_keys = ( self.get_in_memory_keys_to_update() ) # if no pattern OR redis cache does not support scan_iter, use in-memory keys 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 b47a2f1c90f..5696689f1eb 100644 --- a/tests/test_litellm/router_strategy/test_base_routing_strategy.py +++ b/tests/test_litellm/router_strategy/test_base_routing_strategy.py @@ -1,6 +1,7 @@ import json import os import sys +from typing import Any, Dict, List, Optional, Set, Union import pytest @@ -25,20 +26,20 @@ def mock_dual_cache(): dual_cache.redis_cache = MagicMock() # Set up async method mocks to return coroutines - future1 = asyncio.Future() + future1: asyncio.Future[None] = asyncio.Future() future1.set_result(None) dual_cache.in_memory_cache.async_increment.return_value = future1 - future2 = asyncio.Future() + future2: asyncio.Future[None] = asyncio.Future() future2.set_result(None) dual_cache.redis_cache.async_increment_pipeline.return_value = future2 - future3 = asyncio.Future() + future3: asyncio.Future[None] = 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: asyncio.Future[Dict[str, str]] = asyncio.Future() batch_future.set_result({"key1": "10.0", "key2": "20.0"}) dual_cache.redis_cache.async_batch_get_cache.return_value = batch_future @@ -96,23 +97,48 @@ async def test_push_in_memory_increments_to_redis(base_strategy, mock_dual_cache @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"} + base_strategy.in_memory_keys_to_update = {"key1"} + + # Mock the in-memory cache batch get responses + 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 Redis batch get response + redis_future: asyncio.Future[Dict[str, str]] = asyncio.Future() + redis_future.set_result({"key1": "15.0"}) # Redis values + mock_dual_cache.redis_cache.async_batch_get_cache.return_value = redis_future + + # Mock in-memory after snapshot + in_memory_after_future: asyncio.Future[List[str]] = asyncio.Future() + in_memory_after_future.set_result(["8.0"]) # Values after potential updates + mock_dual_cache.in_memory_cache.async_batch_get_cache.side_effect = [ + in_memory_before_future, # First call for before snapshot + in_memory_after_future, # Second call for after snapshot + ] - # 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 + # Verify Redis batch get was called with correct keys key_list = mock_dual_cache.redis_cache.async_batch_get_cache.call_args.kwargs[ "key_list" ] + assert sorted(key_list) == sorted(["key1"]) - 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 with merged values + # For key1: redis_val(15.0) + delta(8.0 - 5.0) = 18.0 + # For key2: redis_val(20.0) + delta(12.0 - 10.0) = 22.0 + assert mock_dual_cache.in_memory_cache.async_set_cache.call_count == 1 - # Verify in-memory cache was updated - assert mock_dual_cache.in_memory_cache.async_set_cache.call_count == 2 + # 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 call.kwargs["value"] == 18.0 + for call in set_cache_calls + ) # Verify cache keys were reset assert len(base_strategy.in_memory_keys_to_update) == 0