fix(proxy): avoid auth cache rewrite during spend updates

This commit is contained in:
Aditya Pratap Singh Hada 2026-07-06 19:39:10 +05:30
parent 5b93ba0ada
commit 5e8ef93d78
2 changed files with 43 additions and 15 deletions

View file

@ -2779,21 +2779,6 @@ async def update_cache(
)
# set cooldown on alert
if existing_spend_obj is not None and getattr(existing_spend_obj, "team_spend", None) is not None:
existing_team_spend = existing_spend_obj.team_spend or 0
# Calculate the new cost by adding the existing cost and response_cost
existing_spend_obj.team_spend = existing_team_spend + response_cost
if existing_spend_obj is not None and getattr(existing_spend_obj, "team_member_spend", None) is not None:
existing_team_member_spend = existing_spend_obj.team_member_spend or 0
# Calculate the new cost by adding the existing cost and response_cost
existing_spend_obj.team_member_spend = existing_team_member_spend + response_cost
# Existing spend_obj is mutated; UserApiKeyCache.async_set_cache_pipeline turns
# BaseModel values into dicts for Redis (same Codec path as async_set_cache).
existing_spend_obj.spend = new_spend
values_to_update_in_cache.append((hashed_token, existing_spend_obj))
### UPDATE USER SPEND ###
async def _update_user_cache():
## UPDATE CACHE FOR USER ID + GLOBAL PROXY

View file

@ -27,6 +27,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import UserAPIKeyAuth
from .conftest import normalize
@ -1250,6 +1251,48 @@ async def test_update_cache_no_cached_entities_schedules_pipeline_flush(monkeypa
}
@pytest.mark.asyncio
async def test_update_cache_does_not_rewrite_cached_key_auth_object(monkeypatch):
cached_key = UserAPIKeyAuth(
token="hashed-token",
models=["old-model"],
spend=10.0,
team_spend=20.0,
team_member_spend=30.0,
)
fake_user_cache = _make_user_api_key_cache(get_value=cached_key)
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
await ps.update_cache(
token="hashed-token",
user_id=None,
end_user_id=None,
team_id=None,
response_cost=1.0,
parent_otel_span=None,
tags=None,
)
await asyncio.sleep(0.1)
cache_list = fake_user_cache.async_set_cache_pipeline.call_args.kwargs[
"cache_list"
]
observed = {
"cache_list": cache_list,
"spend": cached_key.spend,
"team_spend": cached_key.team_spend,
"team_member_spend": cached_key.team_member_spend,
"models": cached_key.models,
}
assert normalize(observed) == {
"cache_list": [],
"spend": 10.0,
"team_spend": 20.0,
"team_member_spend": 30.0,
"models": ["old-model"],
}
@pytest.mark.asyncio
async def test_update_cache_user_cache_failure_invalid_state_is_swallowed(monkeypatch):
"""An inner _update_user_cache raising must not propagate — update_cache