diff --git a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py index c25d8533128..67a24567461 100644 --- a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py +++ b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py @@ -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 diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index c6b830371ba..914be364579 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -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 + ) diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index e035168ca00..50339210a6e 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -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, ): """ diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d2dfa177515..8c5fce84099 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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.