diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 40c8caa49e5..728307cbec2 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -74,10 +74,16 @@ class ResetBudgetJob: except Exception as redis_err: verbose_proxy_logger.warning( "Failed to reset spend counter %s in Redis: %s. " - "Budget may be over-enforced until counter expires.", + "Falling back to DELETE to force reseed from DB on next read.", counter_key, redis_err, ) + try: + await spend_counter_cache.redis_cache.async_delete_cache( + key=counter_key + ) + except Exception: + pass except Exception as e: verbose_proxy_logger.warning( "Failed to reset spend counter %s: %s", counter_key, e diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 80ded0bdd16..c12b51184a8 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -4730,6 +4730,22 @@ async def reset_key_spend_fn( proxy_logging_obj=proxy_logging_obj, ) + from litellm.proxy.proxy_server import spend_counter_cache + + counter_key = f"spend:key:{hashed_api_key}" + spend_counter_cache.in_memory_cache.delete_cache(key=counter_key) + if spend_counter_cache.redis_cache is not None: + try: + await spend_counter_cache.redis_cache.async_delete_cache( + key=counter_key + ) + except Exception as redis_err: + verbose_proxy_logger.warning( + "Failed to delete Redis spend counter %s: %s", + counter_key, + redis_err, + ) + max_budget = updated_key.max_budget budget_reset_at = updated_key.budget_reset_at diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 0b683745369..9dce095386d 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -1336,6 +1336,7 @@ def _make_counter_invalidation_job(monkeypatch): spend_counter_cache.in_memory_cache.set_cache = MagicMock() spend_counter_cache.redis_cache = MagicMock() spend_counter_cache.redis_cache.async_set_cache = AsyncMock() + spend_counter_cache.redis_cache.async_delete_cache = AsyncMock() user_api_key_cache = MagicMock() user_api_key_cache.async_delete_cache = AsyncMock() @@ -1803,3 +1804,38 @@ def test_reset_budget_for_tags_linked_to_budgets_management_cache_delete_failure asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) prisma_client.db.litellm_tagtable.update_many.assert_awaited_once() + + +def test_invalidate_spend_counter_deletes_on_redis_set_failure(monkeypatch): + """When Redis SET fails, _invalidate_spend_counter must fall back to DELETE.""" + counter_cache = _make_counter_invalidation_job(monkeypatch) + counter_cache.redis_cache.async_set_cache = AsyncMock( + side_effect=RuntimeError("SET failed") + ) + + asyncio.run( + ResetBudgetJob._invalidate_spend_counter("spend:key:test-fallback") + ) + + counter_cache.redis_cache.async_delete_cache.assert_awaited_once_with( + key="spend:key:test-fallback" + ) + + +def test_invalidate_spend_counter_swallows_delete_failure(monkeypatch): + """When both Redis SET and DELETE fallback fail, the method must not raise.""" + counter_cache = _make_counter_invalidation_job(monkeypatch) + counter_cache.redis_cache.async_set_cache = AsyncMock( + side_effect=RuntimeError("SET failed") + ) + counter_cache.redis_cache.async_delete_cache = AsyncMock( + side_effect=RuntimeError("DELETE also failed") + ) + + asyncio.run( + ResetBudgetJob._invalidate_spend_counter("spend:key:test-double-fail") + ) + + counter_cache.redis_cache.async_delete_cache.assert_awaited_once_with( + key="spend:key:test-double-fail" + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 8fb242372e3..534b92b6043 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -1,6 +1,7 @@ import json import os import sys +import types import litellm import pytest @@ -6279,6 +6280,155 @@ async def test_reset_key_spend_success(monkeypatch): mock_delete_cache.assert_awaited_once() +@pytest.mark.asyncio +async def test_reset_key_spend_invalidates_redis_spend_counter(monkeypatch): + mock_prisma_client = MagicMock() + mock_user_api_key_cache = MagicMock() + mock_proxy_logging_obj = MagicMock() + + hashed_key = "hashed-test-key" + key_in_db = LiteLLM_VerificationToken( + token=hashed_key, + user_id="test-user", + spend=100.0, + max_budget=200.0, + litellm_budget_table=None, + ) + updated_key = LiteLLM_VerificationToken( + token=hashed_key, + user_id="test-user", + spend=50.0, + max_budget=200.0, + budget_reset_at=None, + ) + + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=key_in_db + ) + mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock( + return_value=updated_key + ) + + spend_counter_cache = MagicMock() + spend_counter_cache.in_memory_cache.delete_cache = MagicMock() + spend_counter_cache.redis_cache = MagicMock() + spend_counter_cache.redis_cache.async_delete_cache = AsyncMock() + + fake_module = types.ModuleType("litellm.proxy.proxy_server") + fake_module.prisma_client = mock_prisma_client + fake_module.hash_token = MagicMock(return_value=hashed_key) + fake_module.user_api_key_cache = mock_user_api_key_cache + fake_module.proxy_logging_obj = mock_proxy_logging_obj + fake_module.spend_counter_cache = spend_counter_cache + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module) + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" + ) as mock_check_admin, + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + ) as mock_delete_cache, + ): + mock_check_admin.return_value = None + mock_delete_cache.return_value = None + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin-user", + ) + + response = await reset_key_spend_fn( + key="sk-test-key", + data=ResetSpendRequest(reset_to=50.0), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert response["spend"] == 50.0 + + counter_key = f"spend:key:{hashed_key}" + spend_counter_cache.in_memory_cache.delete_cache.assert_called_once_with( + key=counter_key + ) + spend_counter_cache.redis_cache.async_delete_cache.assert_awaited_once_with( + key=counter_key + ) + + +@pytest.mark.asyncio +async def test_reset_key_spend_redis_delete_failure_does_not_raise(monkeypatch): + """When Redis delete fails, reset_key_spend_fn logs a warning and still succeeds.""" + mock_prisma_client = MagicMock() + mock_user_api_key_cache = MagicMock() + mock_proxy_logging_obj = MagicMock() + + hashed_key = "hashed-test-key" + key_in_db = LiteLLM_VerificationToken( + token=hashed_key, + user_id="test-user", + spend=100.0, + max_budget=200.0, + litellm_budget_table=None, + ) + updated_key = LiteLLM_VerificationToken( + token=hashed_key, + user_id="test-user", + spend=50.0, + max_budget=200.0, + budget_reset_at=None, + ) + + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=key_in_db + ) + mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock( + return_value=updated_key + ) + + spend_counter_cache = MagicMock() + spend_counter_cache.in_memory_cache.delete_cache = MagicMock() + spend_counter_cache.redis_cache = MagicMock() + spend_counter_cache.redis_cache.async_delete_cache = AsyncMock( + side_effect=RuntimeError("redis unavailable") + ) + + fake_module = types.ModuleType("litellm.proxy.proxy_server") + fake_module.prisma_client = mock_prisma_client + fake_module.hash_token = MagicMock(return_value=hashed_key) + fake_module.user_api_key_cache = mock_user_api_key_cache + fake_module.proxy_logging_obj = mock_proxy_logging_obj + fake_module.spend_counter_cache = spend_counter_cache + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module) + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" + ) as mock_check_admin, + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + ) as mock_delete_cache, + ): + mock_check_admin.return_value = None + mock_delete_cache.return_value = None + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin-user", + ) + + response = await reset_key_spend_fn( + key="sk-test-key", + data=ResetSpendRequest(reset_to=50.0), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert response["spend"] == 50.0 + + @pytest.mark.asyncio async def test_reset_key_spend_success_team_admin(monkeypatch): """Test that team admin can reset key spend for keys in their team."""