From 4bd64c872a07dc600bd0f6e504d5f047f750adcb Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Tue, 20 May 2025 23:08:26 -0700 Subject: [PATCH] =?UTF-8?q?fix(internal=5Fuser=5Fendpoints.py):=20allow=20?= =?UTF-8?q?resetting=20spend/max=20budget=20on=20=E2=80=A6=20(#10993)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(internal_user_endpoints.py): allow resetting spend/max budget on user update Fixes https://github.com/BerriAI/litellm/issues/10495 * fix(internal_user_endpoints.py): correctly return set spend for user on /user/new * fix(auth_checks.py): check redis for key object before checking in-memory allows for quicker updates * feat(internal_user_endpoints.py): update cache object when user is updated + check redis on user values being updated * fix(auth_checks.py): use redis cache when user updated * fix: set default value of 'expires' to None --- litellm/caching/dual_cache.py | 5 +- .../proxy/_experimental/out/onboarding.html | 1 - litellm/proxy/_new_secret_config.yaml | 3 + litellm/proxy/_types.py | 2 +- litellm/proxy/auth/auth_checks.py | 138 ++++++++++++++++-- .../internal_user_endpoints.py | 73 ++++----- litellm/proxy/utils.py | 2 + tests/litellm/proxy/auth/test_auth_checks.py | 50 +++++++ .../test_internal_user_endpoints.py | 70 ++++++++- 9 files changed, 292 insertions(+), 52 deletions(-) delete mode 100644 litellm/proxy/_experimental/out/onboarding.html diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 8bef3337587..91ce58162f2 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -196,6 +196,7 @@ class DualCache(BaseCache): key, parent_otel_span: Optional[Span] = None, local_only: bool = False, + redis_only: bool = False, **kwargs, ): # Try to fetch from in-memory cache first @@ -204,7 +205,7 @@ class DualCache(BaseCache): f"async get cache: cache key: {key}; local_only: {local_only}" ) result = None - if self.in_memory_cache is not None: + if self.in_memory_cache is not None and not redis_only: in_memory_result = await self.in_memory_cache.async_get_cache( key, **kwargs ) @@ -213,7 +214,7 @@ class DualCache(BaseCache): if in_memory_result is not None: result = in_memory_result - if result is None and self.redis_cache is not None and local_only is False: + if result is None and self.redis_cache is not None and not local_only: # If not found in in-memory cache, try fetching from Redis redis_result = await self.redis_cache.async_get_cache( key, parent_otel_span=parent_otel_span diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index 6bd2ec7f07f..00000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index a6fb25f3a79..74f81459b88 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -71,3 +71,6 @@ model_list: model: mistral/* api_key: os.environ/MISTRAL_API_KEY access_groups: ["beta-models"] + +litellm_settings: + cache: true \ No newline at end of file diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0c118898f0f..088d1dd7d4c 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -702,7 +702,7 @@ class GenerateKeyRequest(KeyRequestBase): class GenerateKeyResponse(KeyRequestBase): key: str # type: ignore key_name: Optional[str] = None - expires: Optional[datetime] + expires: Optional[datetime] = None user_id: Optional[str] = None token_id: Optional[str] = None litellm_budget_table: Optional[Any] = None diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 48e9787d00d..3c759e839ec 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -11,7 +11,8 @@ Run checks for: import asyncio import re import time -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast +from datetime import datetime +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, Union, cast from fastapi import Request, status from pydantic import BaseModel @@ -33,6 +34,7 @@ from litellm.proxy._types import ( LiteLLM_TeamTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, + LiteLLM_VerificationToken, LiteLLMRoutes, LitellmUserRoles, ProxyErrorTypes, @@ -43,7 +45,12 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.route_llm_request import route_request -from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics +from litellm.proxy.utils import ( + InternalUsageCache, + PrismaClient, + ProxyLogging, + log_db_metrics, +) from litellm.router import Router from litellm.utils import get_utc_datetime @@ -640,6 +647,69 @@ async def _get_fuzzy_user_object( return response +class UserObjectCache: + def __init__( + self, + user_api_key_cache: DualCache, + internal_usage_cache: Optional[InternalUsageCache] = None, + ): + """ + - user_api_key_cache: cache for user api keys + - internal_usage_cache: cache for internal usage (connected to Redis) + """ + self.user_api_key_cache = user_api_key_cache + self.internal_usage_cache = internal_usage_cache + + async def update_user_object( + self, + user_id: str, + user_object: Union[dict, LiteLLM_UserTable], + litellm_parent_otel_span: Optional[Span] = None, + ): + """ + - update user object in cache + """ + if isinstance(user_object, LiteLLM_UserTable): + user_object = user_object.model_dump() + for k, v in user_object.items(): + if isinstance(v, datetime): + user_object[k] = v.isoformat() + await self.user_api_key_cache.async_set_cache(key=user_id, value=user_object) + if self.internal_usage_cache is not None: + await self.internal_usage_cache.async_set_cache( + key=user_id, + value=user_object, + litellm_parent_otel_span=litellm_parent_otel_span, + ) + + async def get_user_object( + self, user_id: str, litellm_parent_otel_span: Optional[Span] = None + ) -> Optional[LiteLLM_UserTable]: + """ + - get user object from cache + """ + cached_obj: Optional[Union[dict, LiteLLM_UserTable]] = None + + ## CHECK REDIS CACHE ## + if self.internal_usage_cache is not None: + cached_obj = await self.internal_usage_cache.async_get_cache( + key=user_id, + litellm_parent_otel_span=litellm_parent_otel_span, + redis_only=True, + ) + + if cached_obj is None: + cached_obj = await self.user_api_key_cache.async_get_cache(key=user_id) + + if cached_obj is not None: + if isinstance(cached_obj, dict): + return LiteLLM_UserTable(**cached_obj) + elif isinstance(cached_obj, LiteLLM_UserTable): + return cached_obj + + return None + + @log_db_metrics async def get_user_object( user_id: Optional[str], @@ -657,18 +727,23 @@ async def get_user_object( - if valid, return LiteLLM_UserTable object with defined limits - if not, then raise an error """ + user_object_cache = UserObjectCache( + user_api_key_cache=user_api_key_cache, + internal_usage_cache=proxy_logging_obj.internal_usage_cache + if proxy_logging_obj is not None + else None, + ) if user_id is None: return None # check if in cache if not check_db_only: - cached_user_obj = await user_api_key_cache.async_get_cache(key=user_id) + cached_user_obj = await user_object_cache.get_user_object( + user_id=user_id, litellm_parent_otel_span=parent_otel_span + ) if cached_user_obj is not None: - if isinstance(cached_user_obj, dict): - return LiteLLM_UserTable(**cached_user_obj) - elif isinstance(cached_user_obj, LiteLLM_UserTable): - return cached_user_obj + return cached_user_obj # else, check db if prisma_client is None: raise Exception("No db connected") @@ -726,7 +801,9 @@ async def get_user_object( response_dict = _response.model_dump() # save the user object to cache - await user_api_key_cache.async_set_cache(key=user_id, value=response_dict) + await user_object_cache.update_user_object( + user_id=user_id, user_object=response_dict + ) # save to db access time _update_last_db_access_time( @@ -1020,6 +1097,38 @@ class ExperimentalUIJWTToken: ) +async def _get_object_from_cache( + key: str, + proxy_logging_obj: Optional[ProxyLogging], + user_api_key_cache: DualCache, + parent_otel_span: Optional[Span], + base_model: Type[BaseModel], +) -> Optional[BaseModel]: + cached_obj: Optional[Union[dict, BaseModel]] = None + + ## CHECK REDIS CACHE ## + if ( + proxy_logging_obj is not None + and proxy_logging_obj.internal_usage_cache.dual_cache + ): + cached_obj = ( + await proxy_logging_obj.internal_usage_cache.dual_cache.async_get_cache( + key=key, parent_otel_span=parent_otel_span + ) + ) + + if cached_obj is None: + cached_obj = await user_api_key_cache.async_get_cache(key=key) + + if cached_obj is not None: + if isinstance(cached_obj, dict): + return base_model(**cached_obj) + elif isinstance(cached_obj, base_model): + return cached_obj + + return None + + @log_db_metrics async def get_key_object( hashed_token: str, @@ -1042,15 +1151,16 @@ async def get_key_object( # check if in cache key = hashed_token - cached_key_obj: Optional[UserAPIKeyAuth] = await user_api_key_cache.async_get_cache( - key=key + cached_key_obj = await _get_object_from_cache( + key=key, + proxy_logging_obj=proxy_logging_obj, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + base_model=LiteLLM_VerificationToken, ) if cached_key_obj is not None: - if isinstance(cached_key_obj, dict): - return UserAPIKeyAuth(**cached_key_obj) - elif isinstance(cached_key_obj, UserAPIKeyAuth): - return cached_key_obj + return UserAPIKeyAuth(**cached_key_obj.model_dump(exclude_none=True)) if check_cache_only: raise Exception( diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index c00d282a810..c0d0186c67d 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -23,6 +23,7 @@ from fastapi import APIRouter, Depends, Header, HTTPException, Request, status import litellm from litellm._logging import verbose_proxy_logger from litellm.proxy._types import * +from litellm.proxy.auth.auth_checks import UserObjectCache from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity @@ -54,9 +55,9 @@ def _update_internal_new_user_params(data_json: dict, data: NewUserRequest) -> d data_json["user_id"] = str(uuid.uuid4()) auto_create_key = data_json.pop("auto_create_key", True) if auto_create_key is False: - data_json["table_name"] = ( - "user" # only create a user, don't create key if 'auto_create_key' set to False - ) + data_json[ + "table_name" + ] = "user" # only create a user, don't create key if 'auto_create_key' set to False is_internal_user = False if data.user_role and data.user_role.is_internal_user_role: @@ -238,23 +239,15 @@ async def new_user( else: raise e - new_user_response = NewUserResponse( - key=response.get("token", ""), - expires=response.get("expires", None), - max_budget=response["max_budget"], - user_id=response["user_id"], - user_role=response.get("user_role", None), - user_email=response.get("user_email", None), - user_alias=response.get("user_alias", None), - teams=response.get("teams", None), - team_id=response.get("team_id", None), - metadata=response.get("metadata", None), - models=response.get("models", None), - tpm_limit=response.get("tpm_limit", None), - rpm_limit=response.get("rpm_limit", None), - budget_duration=response.get("budget_duration", None), - model_max_budget=response.get("model_max_budget", None), - ) + special_keys = ["token"] + response_dict = {} + for key, value in response.items(): + if key in NewUserResponse.model_fields.keys() and key not in special_keys: + response_dict[key] = value + + response_dict["key"] = response.get("token", "") + + new_user_response = NewUserResponse(**response_dict) ######################################################### ########## USER CREATED HOOK ################ @@ -555,7 +548,6 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest) -> di not in ( [], {}, - 0, ) and k not in LiteLLM_ManagementEndpoint_MetadataFields ): # models default to [], spend defaults to 0, we should not reset these values @@ -582,9 +574,9 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest) -> di "budget_duration" not in non_default_values ): # applies internal user limits, if user role updated if is_internal_user and litellm.internal_user_budget_duration is not None: - non_default_values["budget_duration"] = ( - litellm.internal_user_budget_duration - ) + non_default_values[ + "budget_duration" + ] = litellm.internal_user_budget_duration from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time non_default_values["budget_reset_at"] = get_budget_reset_time( @@ -650,10 +642,15 @@ async def user_update( """ - from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client + from litellm.proxy.proxy_server import ( + litellm_proxy_admin_name, + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) try: - data_json: dict = data.json() + data_json: dict = data.model_dump(exclude_unset=True) # get the row from db if prisma_client is None: raise Exception("Not connected to DB!") @@ -734,6 +731,16 @@ async def user_update( user_row_litellm_typed = LiteLLM_UserTable( **user_row.model_dump(exclude_none=True) ) + + ## UPDATE CACHE ## + user_object_cache = UserObjectCache( + user_api_key_cache=user_api_key_cache, + internal_usage_cache=proxy_logging_obj.internal_usage_cache, + ) + await user_object_cache.update_user_object( + user_id=response["user_id"], user_object=user_row_litellm_typed + ) + asyncio.create_task( UserManagementEventHooks.create_internal_user_audit_log( user_id=user_row_litellm_typed.user_id, @@ -1227,13 +1234,13 @@ async def ui_view_users( } # Query users with pagination and filters - users: Optional[List[BaseModel]] = ( - await prisma_client.db.litellm_usertable.find_many( - where=where_conditions, - skip=skip, - take=page_size, - order={"created_at": "desc"}, - ) + users: Optional[ + List[BaseModel] + ] = await prisma_client.db.litellm_usertable.find_many( + where=where_conditions, + skip=skip, + take=page_size, + order={"created_at": "desc"}, ) if not users: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d1e725def55..64eb0a5cb46 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -146,12 +146,14 @@ class InternalUsageCache: key, litellm_parent_otel_span: Union[Span, None], local_only: bool = False, + redis_only: bool = False, **kwargs, ) -> Any: return await self.dual_cache.async_get_cache( key=key, local_only=local_only, parent_otel_span=litellm_parent_otel_span, + redis_only=redis_only, **kwargs, ) diff --git a/tests/litellm/proxy/auth/test_auth_checks.py b/tests/litellm/proxy/auth/test_auth_checks.py index 97c952b63f5..aa60952add4 100644 --- a/tests/litellm/proxy/auth/test_auth_checks.py +++ b/tests/litellm/proxy/auth/test_auth_checks.py @@ -13,8 +13,10 @@ from datetime import datetime, timedelta import pytest import litellm +from litellm.caching.dual_cache import DualCache from litellm.proxy._types import ( LiteLLM_UserTable, + LiteLLM_VerificationToken, LitellmUserRoles, SSOUserDefinedValues, ) @@ -115,6 +117,54 @@ def test_get_key_object_from_ui_hash_key_invalid(): @pytest.mark.asyncio +async def test_get_object_from_cache_redis(): + """Test that _get_object_from_cache retrieves from Redis cache when available""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_VerificationToken + from litellm.proxy.auth.auth_checks import _get_object_from_cache + + # Create mock objects + mock_proxy_logging = MagicMock() + mock_proxy_logging.internal_usage_cache.dual_cache = AsyncMock() + mock_user_api_key_cache = DualCache() + + # Create test data + test_key = "test_key" + test_data = { + "token": "test_token", + "key_name": "test_key_name", + "spend": 0.0, + "models": ["gpt-3.5-turbo"], + } + + # Mock Redis cache response + mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache.return_value = ( + test_data + ) + + # Call the function + result = await _get_object_from_cache( + key=test_key, + proxy_logging_obj=mock_proxy_logging, + user_api_key_cache=mock_user_api_key_cache, + parent_otel_span=None, + base_model=LiteLLM_VerificationToken, + ) + + # Verify Redis cache was checked + mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache.assert_called_once_with( + key=test_key, parent_otel_span=None + ) + + # Verify result is correct + assert isinstance(result, LiteLLM_VerificationToken) + assert result.token == "test_token" + assert result.key_name == "test_key_name" + assert result.spend == 0.0 + assert result.models == ["gpt-3.5-turbo"] + async def test_default_internal_user_params_with_get_user_object(monkeypatch): """Test that default_internal_user_params is used when creating a new user via get_user_object""" # Set up default_internal_user_params diff --git a/tests/litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 360f21f1717..53653a2933c 100644 --- a/tests/litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -10,9 +10,14 @@ sys.path.insert( 0, os.path.abspath("../../../..") ) # Adds the parent directory to the system path -from litellm.proxy._types import LiteLLM_UserTableFiltered, UserAPIKeyAuth +from litellm.proxy._types import ( + LiteLLM_UserTableFiltered, + UpdateUserRequest, + UserAPIKeyAuth, +) from litellm.proxy.management_endpoints.internal_user_endpoints import ( LiteLLM_UserTableWithKeyCount, + _update_internal_user_params, get_user_key_counts, get_users, ui_view_users, @@ -169,3 +174,66 @@ def test_validate_sort_params(): assert _validate_sort_params("user_id", "desc") == {"user_id": "desc"} with pytest.raises(Exception): _validate_sort_params("user_id", "invalid") + + +def test_update_user_request_pydantic_object(): + """ + Test that _update_internal_user_params correctly processes an email-only update + """ + data = UpdateUserRequest(user_email="test@example.com") + + data_json = data.model_dump(exclude_unset=True) + + assert data_json == {"user_email": "test@example.com"} + + +def test_update_internal_user_params_email(): + """ + Test that _update_internal_user_params correctly processes an email-only update + """ + from litellm.proxy._types import UpdateUserRequest + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_internal_user_params, + ) + + # Create test data with only email update + data_json = {"user_email": "test@example.com"} + data = UpdateUserRequest(user_email="test@example.com") + + # Call the function + non_default_values = _update_internal_user_params(data_json=data_json, data=data) + + # Assertions + assert len(non_default_values) == 1 # Should only contain email + assert "user_email" in non_default_values + assert non_default_values["user_email"] == "test@example.com" + assert "user_id" not in non_default_values # Should not add user_id if not provided + assert "max_budget" not in non_default_values # Should not add default values + assert "budget_duration" not in non_default_values # Should not add default values + + +def test_update_internal_user_params_reset_spend_and_max_budget(): + """ + Relevant Issue: https://github.com/BerriAI/litellm/issues/10495 + """ + from litellm.proxy._types import UpdateUserRequest + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_internal_user_params, + ) + + # Create test data with only email update + data = UpdateUserRequest(spend=0, max_budget=0, user_id="test_user_id") + data_json = data.model_dump(exclude_unset=True) + + # Call the function + non_default_values = _update_internal_user_params(data_json=data_json, data=data) + + # Assertions + assert len(non_default_values) == 3 # Should only contain email + assert "spend" in non_default_values + assert non_default_values["spend"] == 0 + assert "max_budget" in non_default_values + assert non_default_values["max_budget"] == 0 + assert "user_id" in non_default_values # Should not add user_id if not provided + assert non_default_values["user_id"] == "test_user_id" + assert "budget_duration" not in non_default_values # Should not add default values