fix: preserve model budget spend on key regeneration

This commit is contained in:
Aditya Pratap Singh Hada 2026-06-23 17:28:26 +05:30
parent 69b0dd2da0
commit f691b4ab54
2 changed files with 278 additions and 0 deletions

View file

@ -17,6 +17,7 @@ import math
import os
import re
import secrets
import time
import traceback
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
@ -3289,6 +3290,171 @@ async def _build_model_max_budget_usage(
return result
VIRTUAL_KEY_BUDGET_START_TIME_CACHE_KEY_PREFIX = "virtual_key_budget_start_time"
def _coerce_model_max_budget(value: object) -> Mapping[str, object]:
if value is None:
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except Exception: # noqa: BLE001
return {}
if isinstance(value, Mapping):
return value
return {}
def _get_model_max_budget_for_cache_migration(
key_in_db: LiteLLM_VerificationToken,
non_default_values: dict,
) -> Mapping[str, object]:
if "model_max_budget" in non_default_values:
return _coerce_model_max_budget(non_default_values.get("model_max_budget"))
return _coerce_model_max_budget(getattr(key_in_db, "model_max_budget", None))
def _model_max_budget_spend_cache_keys(
api_key_hash: str, model: str, budget_duration: str
) -> tuple[str, ...]:
keys = [
f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:"
f"{api_key_hash}:{model}:{budget_duration}"
]
if "/" in model:
model_without_prefix = model.split("/")[-1]
keys.append(
f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:"
f"{api_key_hash}:{model_without_prefix}:{budget_duration}"
)
return tuple(keys)
def _normalize_cache_ttl_seconds(ttl: object) -> int | None:
if ttl is None:
return None
try:
ttl_seconds = int(float(ttl))
except (TypeError, ValueError):
return None
if ttl_seconds > time.time():
ttl_seconds = int(ttl_seconds - time.time())
if ttl_seconds <= 0:
return None
return ttl_seconds
def _remaining_ttl_from_budget_start_time(
budget_start_time: object, budget_duration_seconds: int
) -> int | None:
try:
remaining_ttl = budget_duration_seconds - (
time.time() - float(budget_start_time)
)
except (TypeError, ValueError):
return None
if remaining_ttl <= 0:
return None
return int(remaining_ttl)
async def _get_cache_ttl_seconds(
cache_key: str,
user_api_key_cache: UserApiKeyCache,
) -> int | None:
ttl = await user_api_key_cache.async_get_ttl(cache_key)
return _normalize_cache_ttl_seconds(ttl)
async def _copy_cache_value(
source_key: str,
destination_key: str,
ttl: int | None,
user_api_key_cache: UserApiKeyCache,
) -> None:
value = await user_api_key_cache.async_get_cache(key=source_key)
if value is None:
return
if ttl is None:
await user_api_key_cache.async_set_cache(key=destination_key, value=value)
return
await user_api_key_cache.async_set_cache(
key=destination_key,
value=value,
ttl=ttl,
)
async def _migrate_model_max_budget_cache_on_regenerate(
*,
old_api_key_hash: str,
new_api_key_hash: str,
model_max_budget: Mapping[str, object],
user_api_key_cache: UserApiKeyCache,
) -> None:
if not model_max_budget:
return
start_time_source_key = (
f"{VIRTUAL_KEY_BUDGET_START_TIME_CACHE_KEY_PREFIX}:{old_api_key_hash}"
)
start_time_destination_key = (
f"{VIRTUAL_KEY_BUDGET_START_TIME_CACHE_KEY_PREFIX}:{new_api_key_hash}"
)
budget_start_time = await user_api_key_cache.async_get_cache(
key=start_time_source_key
)
for model, budget_info in model_max_budget.items():
try:
budget_config = BudgetConfig.model_validate(budget_info)
if budget_config.budget_duration is None:
continue
budget_duration_seconds = duration_in_seconds(budget_config.budget_duration)
except Exception: # noqa: BLE001
continue
start_time_ttl = await _get_cache_ttl_seconds(
cache_key=start_time_source_key,
user_api_key_cache=user_api_key_cache,
)
if start_time_ttl is None:
start_time_ttl = _remaining_ttl_from_budget_start_time(
budget_start_time=budget_start_time,
budget_duration_seconds=budget_duration_seconds,
)
await _copy_cache_value(
source_key=start_time_source_key,
destination_key=start_time_destination_key,
ttl=start_time_ttl,
user_api_key_cache=user_api_key_cache,
)
old_spend_keys = _model_max_budget_spend_cache_keys(
api_key_hash=old_api_key_hash,
model=model,
budget_duration=budget_config.budget_duration,
)
new_spend_keys = _model_max_budget_spend_cache_keys(
api_key_hash=new_api_key_hash,
model=model,
budget_duration=budget_config.budget_duration,
)
for source_key, destination_key in zip(old_spend_keys, new_spend_keys):
spend_ttl = await _get_cache_ttl_seconds(
cache_key=source_key,
user_api_key_cache=user_api_key_cache,
)
await _copy_cache_value(
source_key=source_key,
destination_key=destination_key,
ttl=spend_ttl or start_time_ttl,
user_api_key_cache=user_api_key_cache,
)
@router.post(
"/v2/key/info",
tags=["key management"],
@ -4468,6 +4634,16 @@ async def _execute_virtual_key_regeneration(
updated_token_dict["key"] = new_token
updated_token_dict["token_id"] = updated_token_dict.pop("token")
await _migrate_model_max_budget_cache_on_regenerate(
old_api_key_hash=hashed_api_key,
new_api_key_hash=new_token_hash,
model_max_budget=_get_model_max_budget_for_cache_migration(
key_in_db=key_in_db,
non_default_values=non_default_values,
),
user_api_key_cache=user_api_key_cache,
)
if hashed_api_key or key:
await _delete_cache_key_object(
hashed_token=_hash_token_if_needed(key),

View file

@ -10195,6 +10195,22 @@ def _make_regenerate_existing_key():
)
class _RegenerateModelBudgetCache:
def __init__(self):
self.cache = {}
self.ttls = {}
async def async_get_cache(self, key, **kwargs):
return self.cache.get(key)
async def async_set_cache(self, key, value, **kwargs):
self.cache[key] = value
self.ttls[key] = kwargs.get("ttl")
async def async_get_ttl(self, key):
return self.ttls.get(key)
@pytest.mark.asyncio
async def test_execute_virtual_key_regeneration_rejects_over_limit_duration():
"""Regenerate must reject durations exceeding upperbound_key_generate_params.duration."""
@ -10251,6 +10267,92 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_duration():
litellm.upperbound_key_generate_params = original
@pytest.mark.asyncio
async def test_execute_virtual_key_regeneration_migrates_model_budget_cache():
from litellm.proxy.management_endpoints.key_management_endpoints import (
_execute_virtual_key_regeneration,
)
old_token_hash = "old-token-hash"
new_token_hash = "new-token-hash"
existing_key = LiteLLM_VerificationToken(
token=old_token_hash,
user_id="user-1",
models=["gpt-4"],
team_id=None,
max_budget=None,
tags=None,
model_max_budget={
"gpt-4": {
"budget_limit": 0.0001,
"time_period": "1d",
}
},
)
mock_prisma_client = _make_regenerate_mock_prisma()
mock_prisma_client.db.litellm_verificationtoken.update.return_value = (
type(
"DictLikeResult",
(),
{
"__iter__": lambda self: iter(
{
"token": new_token_hash,
"key_name": "sk-...ab12",
"user_id": "user-1",
}.items()
)
},
)()
)
cache = _RegenerateModelBudgetCache()
cache.cache[f"virtual_key_budget_start_time:{old_token_hash}"] = 12345.0
cache.cache[f"virtual_key_spend:{old_token_hash}:gpt-4:1d"] = 0.001
cache.ttls[f"virtual_key_budget_start_time:{old_token_hash}"] = 3600
cache.ttls[f"virtual_key_spend:{old_token_hash}:gpt-4:1d"] = 3600
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.get_new_token",
new_callable=AsyncMock,
return_value="sk-newtoken1234ab12",
),
patch(
"litellm.proxy.proxy_server.hash_token",
return_value=new_token_hash,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook",
new_callable=AsyncMock,
),
):
await _execute_virtual_key_regeneration(
prisma_client=mock_prisma_client,
key_in_db=existing_key,
hashed_api_key=old_token_hash,
key=old_token_hash,
data=None,
user_api_key_dict=_make_regenerate_user_api_key_dict(),
litellm_changed_by=None,
user_api_key_cache=cache,
proxy_logging_obj=MagicMock(),
)
assert cache.cache[f"virtual_key_budget_start_time:{new_token_hash}"] == 12345.0
assert cache.ttls[f"virtual_key_budget_start_time:{new_token_hash}"] == 3600
assert cache.cache[f"virtual_key_spend:{new_token_hash}:gpt-4:1d"] == 0.001
assert cache.ttls[f"virtual_key_spend:{new_token_hash}:gpt-4:1d"] == 3600
@pytest.mark.asyncio
async def test_execute_virtual_key_regeneration_allows_within_limit_duration():
"""Regenerate must accept durations within upperbound_key_generate_params.duration."""