mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: preserve model budget spend on key regeneration
This commit is contained in:
parent
69b0dd2da0
commit
f691b4ab54
2 changed files with 278 additions and 0 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue