mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): avoid auth cache rewrite during spend updates
This commit is contained in:
parent
5b93ba0ada
commit
5e8ef93d78
2 changed files with 43 additions and 15 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue