From f691b4ab546b009ddffc4967843779355059ca2b Mon Sep 17 00:00:00 2001 From: Aditya Pratap Singh Hada Date: Tue, 23 Jun 2026 17:28:26 +0530 Subject: [PATCH] fix: preserve model budget spend on key regeneration --- .../key_management_endpoints.py | 176 ++++++++++++++++++ .../test_key_management_endpoints.py | 102 ++++++++++ 2 files changed, 278 insertions(+) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2d49297c8e9..8b1917d179a 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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), 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 b8ec8a8a388..e70c182b8f9 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 @@ -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."""