refactor: replace DualCache with UserApiKeyCache in multiple modules

This commit updates the codebase to replace instances of DualCache with UserApiKeyCache in various files, including utils, expired_ui_session_key_cleanup_manager, and team_member_permission_checks. Additionally, it enhances the UserApiKeyCache class with new methods for cache management, improving type safety and consistency across the application.
This commit is contained in:
harish-berri 2026-04-30 22:40:52 +00:00
parent 3bff192c2f
commit 6fb1de9763
4 changed files with 59 additions and 58 deletions

View file

@ -8,7 +8,7 @@ from datetime import datetime, timezone
from typing import Any, Dict, List, Optional
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.constants import (
EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME,
LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE,
@ -31,7 +31,7 @@ class ExpiredUISessionKeyCleanupManager:
def __init__(
self,
prisma_client: PrismaClient,
user_api_key_cache: DualCache,
user_api_key_cache: UserApiKeyCache,
pod_lock_manager=None,
):
self.prisma_client = prisma_client

View file

@ -31,11 +31,11 @@ class UserApiKeyCache(DualCache):
``async_set_cache_pipeline`` applies the same untyped Codec pass as omitting
``model_type`` on ``async_set_cache`` (so ``BaseModel`` rows are dumped before Redis).
``get_cache`` / ``async_get_cache`` overloads and implementations must be contiguous
(no other methods in between) so mypy resolves ``@overload`` + implementation correctly.
"""
# Overloads: `model_type` must be a real parameter (not only via **kwargs) so
# the untyped branch cannot match calls that pass `model_type=...`.
@overload
def get_cache(
self,
@ -56,55 +56,6 @@ class UserApiKeyCache(DualCache):
**kwargs: Any,
) -> Any: ...
def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override]
model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None))
payload = CacheCodec.serialize(value, model_type=model_type)
return super().set_cache(
key=key, value=payload, local_only=local_only, **kwargs
)
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override]
model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None))
payload = CacheCodec.serialize(value, model_type=model_type)
return await super().async_set_cache(
key=key, value=payload, local_only=local_only, **kwargs
)
async def async_set_cache_pipeline( # type: ignore[override]
self, cache_list: list, local_only: bool = False, **kwargs
) -> None:
"""
Batch writes with the same Codec boundary as ``async_set_cache`` without
``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged.
"""
normalized = [
(key, CacheCodec.serialize(value, model_type=None))
for key, value in cache_list
]
return await super().async_set_cache_pipeline(
cache_list=normalized, local_only=local_only, **kwargs
)
@overload
async def async_get_cache(
self,
key: Any,
parent_otel_span: Any = None,
local_only: bool = False,
*,
model_type: Type[T],
**kwargs: Any,
) -> Optional[T]: ...
@overload
async def async_get_cache(
self,
key: Any,
parent_otel_span: Any = None,
local_only: bool = False,
**kwargs: Any,
) -> Any: ...
def get_cache( # type: ignore[override]
self,
key,
@ -133,6 +84,26 @@ class UserApiKeyCache(DualCache):
return None
return decoded
@overload
async def async_get_cache(
self,
key: Any,
parent_otel_span: Any = None,
local_only: bool = False,
*,
model_type: Type[T],
**kwargs: Any,
) -> Optional[T]: ...
@overload
async def async_get_cache(
self,
key: Any,
parent_otel_span: Any = None,
local_only: bool = False,
**kwargs: Any,
) -> Any: ...
async def async_get_cache( # type: ignore[override]
self,
key,
@ -160,3 +131,32 @@ class UserApiKeyCache(DualCache):
)
return None
return decoded
def set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override]
model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None))
payload = CacheCodec.serialize(value, model_type=model_type)
return super().set_cache(
key=key, value=payload, local_only=local_only, **kwargs
)
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): # type: ignore[override]
model_type = cast(Optional[Type[BaseModel]], kwargs.pop("model_type", None))
payload = CacheCodec.serialize(value, model_type=model_type)
return await super().async_set_cache(
key=key, value=payload, local_only=local_only, **kwargs
)
async def async_set_cache_pipeline( # type: ignore[override]
self, cache_list: list, local_only: bool = False, **kwargs
) -> None:
"""
Batch writes with the same Codec boundary as ``async_set_cache`` without
``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged.
"""
normalized = [
(key, CacheCodec.serialize(value, model_type=None))
for key, value in cache_list
]
return await super().async_set_cache_pipeline(
cache_list=normalized, local_only=local_only, **kwargs
)

View file

@ -1,6 +1,5 @@
from typing import List, Optional
from litellm.caching import DualCache
from litellm.proxy._types import (
KeyManagementRoutes,
LiteLLM_TeamTableCachedObj,
@ -12,6 +11,7 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.utils import PrismaClient
@ -65,7 +65,7 @@ class TeamMemberPermissionChecks:
user_api_key_dict: UserAPIKeyAuth,
route: KeyManagementRoutes,
prisma_client: PrismaClient,
user_api_key_cache: DualCache,
user_api_key_cache: UserApiKeyCache,
existing_key_row: LiteLLM_VerificationToken,
):
"""

View file

@ -101,6 +101,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.db.create_views import (
create_missing_views,
should_create_missing_views,
@ -340,7 +341,7 @@ class ProxyLogging:
def __init__(
self,
user_api_key_cache: DualCache,
user_api_key_cache: UserApiKeyCache,
premium_user: bool = False,
):
## INITIALIZE LITELLM CALLBACKS ##
@ -5715,7 +5716,7 @@ async def get_available_models_for_user(
include_model_access_groups: bool = False,
only_model_access_groups: bool = False,
return_wildcard_routes: bool = False,
user_api_key_cache: Optional["DualCache"] = None,
user_api_key_cache: Optional["UserApiKeyCache"] = None,
) -> List[str]:
"""
Get the list of models available to a user based on their API key and team permissions.