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