fix(proxy): invalidate stale key spend counter after budget reset or manual spend update (#30001)

* fix(proxy): reconcile stale key spend counter after budget reset

* fix(proxy): invalidate stale key spend counter after budget reset or manual spend update

* fix(proxy): remove read-time stale counter reconciliation to prevent budget bypass

* revert: undo unrelated formatting changes in enterprise directory

* test(proxy): add unit test for key spend update invalidating counter

* test(proxy): fix mocked update_data and hash token expectations in unit test
This commit is contained in:
Dimitris Spachos 2026-06-10 13:09:16 +03:00 • committed by Sameer Kankute
parent 435809aac2
commit 7729ff5a13
No known key found for this signature in database
4 changed files with 98 additions and 2 deletions

View file

@ -3516,10 +3516,13 @@ async def _virtual_key_max_budget_check(
if valid_token.max_budget is not None:
from litellm.proxy.proxy_server import get_current_spend
fallback_spend = valid_token.spend or 0.0
counter_key = f"spend:key:{valid_token.token}"
# Read spend from cross-pod counter (Redis-first) or cached object (fallback)
spend = await get_current_spend(
counter_key=f"spend:key:{valid_token.token}",
fallback_spend=valid_token.spend or 0.0,
counter_key=counter_key,
fallback_spend=fallback_spend,
)
####################################

View file

@ -2587,6 +2587,17 @@ async def update_key_fn( # noqa: PLR0915
proxy_logging_obj=proxy_logging_obj,
)
if data.spend is not None:
try:
from litellm.proxy.proxy_server import _invalidate_spend_counter
token_to_invalidate = _hash_token_if_needed(key)
await _invalidate_spend_counter(
counter_key=f"spend:key:{token_to_invalidate}"
)
except Exception:
pass
asyncio.create_task(
KeyManagementEventHooks.async_key_updated_hook(
data=data,
@ -4774,6 +4785,13 @@ async def reset_key_spend_fn(
proxy_logging_obj=proxy_logging_obj,
)
try:
from litellm.proxy.proxy_server import _invalidate_spend_counter
await _invalidate_spend_counter(counter_key=f"spend:key:{hashed_api_key}")
except Exception:
pass
max_budget = updated_key.max_budget
budget_reset_at = updated_key.budget_reset_at

View file

@ -2409,6 +2409,8 @@ async def test_virtual_key_budget_check_fallback_no_counter():
assert exc_info.value.current_cost == 15.0
@pytest.mark.asyncio
async def test_team_budget_check_reads_from_spend_counter():
"""Team budget check should use get_current_spend when counter exists."""

View file

@ -6496,6 +6496,9 @@ async def test_reset_key_spend_success(monkeypatch):
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
) as mock_delete_cache,
patch(
"litellm.proxy.proxy_server._invalidate_spend_counter"
) as mock_invalidate,
):
mock_hash_token.return_value = hashed_key
mock_check_admin.return_value = None
@ -6520,6 +6523,76 @@ async def test_reset_key_spend_success(monkeypatch):
assert response["max_budget"] == 200.0
mock_prisma_client.db.litellm_verificationtoken.update.assert_called_once()
mock_delete_cache.assert_awaited_once()
mock_invalidate.assert_awaited_once_with(counter_key=f"spend:key:{hashed_key}")
@pytest.mark.asyncio
async def test_update_key_spend_invalidates_counter(monkeypatch):
"""
Test that updating a key's spend via update_key_fn immediately invalidates the spend counter.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = AsyncMock()
mock_proxy_logging_obj = MagicMock()
hashed_key = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
key_in_db = LiteLLM_VerificationToken(
token=hashed_key,
user_id="test-user",
spend=10.0,
max_budget=200.0,
litellm_budget_table=None,
)
mock_prisma_client.get_data = AsyncMock(return_value=key_in_db)
mock_prisma_client.update_data = AsyncMock(return_value={"data": {"spend": 0.0}})
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=key_in_db
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
monkeypatch.setattr("litellm.store_audit_logs", False)
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
) as mock_delete_cache,
patch(
"litellm.proxy.proxy_server._invalidate_spend_counter"
) as mock_invalidate,
):
mock_delete_cache.return_value = None
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
mock_request = MagicMock()
mock_request.query_params = {}
await update_key_fn(
request=mock_request,
data=UpdateKeyRequest(key="sk-test-key", spend=0.0),
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
mock_delete_cache.assert_awaited_once()
mock_invalidate.assert_awaited_once_with(counter_key=f"spend:key:{hashed_key}")
@pytest.mark.asyncio