diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 210996a0a86..d3b028f63f9 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -165,7 +165,6 @@ class JWTHandler: return False def get_team_ids_from_jwt(self, token: dict) -> List[str]: - if self.litellm_jwtauth.team_ids_jwt_field is not None: team_ids: Optional[List[str]] = get_nested_value( data=token, @@ -245,7 +244,9 @@ class JWTHandler: team_id = default_value return team_id - def get_team_alias(self, token: dict, default_value: Optional[str]) -> Optional[str]: + def get_team_alias( + self, token: dict, default_value: Optional[str] + ) -> Optional[str]: """ Extract team name/alias from JWT token using the configured team_alias_jwt_field. @@ -538,17 +539,17 @@ class JWTHandler: async def get_oidc_userinfo(self, token: str) -> dict: """ Fetch user information from OIDC UserInfo endpoint. - + This follows the OpenID Connect protocol where an access token is sent to the identity provider's UserInfo endpoint to retrieve user identity information. - + Args: token: The access token to use for authentication - + Returns: dict: User information from the UserInfo endpoint - + Raises: Exception: If UserInfo endpoint is not configured or request fails """ @@ -556,19 +557,21 @@ class JWTHandler: raise Exception( "OIDC UserInfo endpoint not configured. Set 'oidc_userinfo_endpoint' in JWT auth config." ) - + # Check cache first - cache_key = f"oidc_userinfo_{token[:20]}" # Use first 20 chars of token as cache key + cache_key = ( + f"oidc_userinfo_{token[:20]}" # Use first 20 chars of token as cache key + ) cached_userinfo = await self.user_api_key_cache.async_get_cache(cache_key) - + if cached_userinfo is not None: verbose_proxy_logger.debug("Returning cached OIDC UserInfo") return cached_userinfo - + verbose_proxy_logger.debug( f"Calling OIDC UserInfo endpoint: {self.litellm_jwtauth.oidc_userinfo_endpoint}" ) - + try: # Call the UserInfo endpoint with the access token response = await self.http_handler.get( @@ -578,24 +581,24 @@ class JWTHandler: "Accept": "application/json", }, ) - + if response.status_code != 200: raise Exception( f"OIDC UserInfo endpoint returned status {response.status_code}: {response.text}" ) - + userinfo = response.json() verbose_proxy_logger.debug(f"Received OIDC UserInfo: {userinfo}") - + # Cache the userinfo response await self.user_api_key_cache.async_set_cache( key=cache_key, value=userinfo, ttl=self.litellm_jwtauth.oidc_userinfo_cache_ttl, ) - + return userinfo - + except Exception as e: verbose_proxy_logger.error(f"Error fetching OIDC UserInfo: {str(e)}") raise Exception(f"Failed to fetch OIDC UserInfo: {str(e)}") @@ -1032,11 +1035,11 @@ class JWTAuthManager: ) -> Tuple[ Optional[LiteLLM_UserTable], Optional[LiteLLM_OrganizationTable], - Optional[LiteLLM_EndUserTable], + Optional[LiteLLM_EndUserTable], Optional[LiteLLM_TeamMembership], ]: """Get user, org, and end user objects. Also resolves org aliases to IDs if configured.""" - + # Get org object - first try by ID, then by alias org_object: Optional[LiteLLM_OrganizationTable] = None if org_id: @@ -1373,7 +1376,9 @@ class JWTAuthManager: # Get team with model access ## Check if team_id is specified via x-litellm-team-id header all_team_ids = JWTAuthManager.get_all_team_ids(jwt_handler, jwt_valid_token) - specific_team_id = jwt_handler.get_team_id(token=jwt_valid_token, default_value=None) + specific_team_id = jwt_handler.get_team_id( + token=jwt_valid_token, default_value=None + ) if specific_team_id: all_team_ids.add(specific_team_id) @@ -1421,22 +1426,25 @@ class JWTAuthManager: org_alias = jwt_handler.get_org_alias(token=jwt_valid_token, default_value=None) # Get other objects - user_object, org_object, end_user_object, team_membership_object = ( - await JWTAuthManager.get_objects( - user_id=user_id, - user_email=user_email, - org_id=org_id, - end_user_id=end_user_id, - team_id=team_id, - valid_user_email=valid_user_email, - jwt_handler=jwt_handler, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - route=route, - org_alias=org_alias, - ) + ( + user_object, + org_object, + end_user_object, + team_membership_object, + ) = await JWTAuthManager.get_objects( + user_id=user_id, + user_email=user_email, + org_id=org_id, + end_user_id=end_user_id, + team_id=team_id, + valid_user_email=valid_user_email, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + route=route, + org_alias=org_alias, ) # Derive org_id from org_object if resolved by alias diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index d453b721645..5c529bc69d6 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -744,10 +744,13 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 team_rpm_limit=( team_object.rpm_limit if team_object is not None else None ), - team_models=team_object.models if team_object is not None else [], + team_models=team_object.models + if team_object is not None + else [], user_role=( LitellmUserRoles(user_object.user_role) - if user_object is not None and user_object.user_role is not None + if user_object is not None + and user_object.user_role is not None else LitellmUserRoles.INTERNAL_USER ), user_id=user_id, diff --git a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py index 056db954a5f..06a526e13a9 100644 --- a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py +++ b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py @@ -1,14 +1,10 @@ -import asyncio -from typing import List, Optional, Union -from fastapi import APIRouter, Depends, HTTPException, Request -import litellm +from fastapi import APIRouter, Depends, HTTPException from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.utils import PrismaClient, ProxyLogging -from litellm.proxy.auth.auth_checks import _delete_cache_key_object router = APIRouter() + @router.post("/jwt/key/mapping/new", tags=["JWT Key Mapping"]) async def create_jwt_key_mapping( data: CreateJWTKeyMappingRequest, @@ -17,7 +13,9 @@ async def create_jwt_key_mapping( from litellm.proxy.proxy_server import prisma_client, user_api_key_cache if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: - raise HTTPException(status_code=403, detail="Only proxy admins can create JWT key mappings") + raise HTTPException( + status_code=403, detail="Only proxy admins can create JWT key mappings" + ) if prisma_client is None: raise HTTPException(status_code=500, detail="Database not connected") @@ -31,7 +29,7 @@ async def create_jwt_key_mapping( "is_active": data.is_active, } ) - + # Invalidate cache cache_key = f"jwt_key_mapping:{data.jwt_claim_name}:{data.jwt_claim_value}" await user_api_key_cache.async_delete_cache(cache_key) @@ -40,6 +38,7 @@ async def create_jwt_key_mapping( except Exception as e: raise HTTPException(status_code=500, detail=str(e)) + @router.post("/jwt/key/mapping/update", tags=["JWT Key Mapping"]) async def update_jwt_key_mapping( data: UpdateJWTKeyMappingRequest, @@ -48,28 +47,29 @@ async def update_jwt_key_mapping( from litellm.proxy.proxy_server import prisma_client, user_api_key_cache if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: - raise HTTPException(status_code=403, detail="Only proxy admins can update JWT key mappings") + raise HTTPException( + status_code=403, detail="Only proxy admins can update JWT key mappings" + ) if prisma_client is None: raise HTTPException(status_code=500, detail="Database not connected") update_data = data.model_dump(exclude_unset=True, exclude={"mapping_id"}) - + try: # Get old mapping for cache invalidation old_mapping = await prisma_client.db.litellm_jwtkeymapping.find_unique( where={"mapping_id": data.mapping_id} ) - + if old_mapping: cache_key = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}" await user_api_key_cache.async_delete_cache(cache_key) updated_mapping = await prisma_client.db.litellm_jwtkeymapping.update( - where={"mapping_id": data.mapping_id}, - data=update_data + where={"mapping_id": data.mapping_id}, data=update_data ) - + # Invalidate new cache key if claim fields changed cache_key = f"jwt_key_mapping:{updated_mapping.jwt_claim_name}:{updated_mapping.jwt_claim_value}" await user_api_key_cache.async_delete_cache(cache_key) @@ -78,6 +78,7 @@ async def update_jwt_key_mapping( except Exception as e: raise HTTPException(status_code=500, detail=str(e)) + @router.post("/jwt/key/mapping/delete", tags=["JWT Key Mapping"]) async def delete_jwt_key_mapping( data: DeleteJWTKeyMappingRequest, @@ -86,7 +87,9 @@ async def delete_jwt_key_mapping( from litellm.proxy.proxy_server import prisma_client, user_api_key_cache if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: - raise HTTPException(status_code=403, detail="Only proxy admins can delete JWT key mappings") + raise HTTPException( + status_code=403, detail="Only proxy admins can delete JWT key mappings" + ) if prisma_client is None: raise HTTPException(status_code=500, detail="Database not connected") @@ -96,7 +99,7 @@ async def delete_jwt_key_mapping( old_mapping = await prisma_client.db.litellm_jwtkeymapping.find_unique( where={"mapping_id": data.mapping_id} ) - + if old_mapping: cache_key = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}" await user_api_key_cache.async_delete_cache(cache_key) @@ -108,6 +111,7 @@ async def delete_jwt_key_mapping( except Exception as e: raise HTTPException(status_code=500, detail=str(e)) + @router.get("/jwt/key/mapping/list", tags=["JWT Key Mapping"]) async def list_jwt_key_mappings( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -115,7 +119,9 @@ async def list_jwt_key_mappings( from litellm.proxy.proxy_server import prisma_client if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: - raise HTTPException(status_code=403, detail="Only proxy admins can list JWT key mappings") + raise HTTPException( + status_code=403, detail="Only proxy admins can list JWT key mappings" + ) if prisma_client is None: raise HTTPException(status_code=500, detail="Database not connected") @@ -126,6 +132,7 @@ async def list_jwt_key_mappings( except Exception as e: raise HTTPException(status_code=500, detail=str(e)) + @router.get("/jwt/key/mapping/info", tags=["JWT Key Mapping"]) async def info_jwt_key_mapping( mapping_id: str, @@ -134,7 +141,9 @@ async def info_jwt_key_mapping( from litellm.proxy.proxy_server import prisma_client if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: - raise HTTPException(status_code=403, detail="Only proxy admins can get JWT key mapping info") + raise HTTPException( + status_code=403, detail="Only proxy admins can get JWT key mapping info" + ) if prisma_client is None: raise HTTPException(status_code=500, detail="Database not connected") diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4c613a4dcbf..d80cb6be577 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2165,22 +2165,24 @@ async def _run_background_health_check(): "Error in shared health check, falling back to direct health check: %s", str(e), ) - healthy_endpoints, unhealthy_endpoints = ( - await _run_direct_health_check_with_instrumentation( - _llm_model_list, - health_check_details, - health_check_concurrency, - instrumentation_context, - ) - ) - else: - healthy_endpoints, unhealthy_endpoints = ( - await _run_direct_health_check_with_instrumentation( + ( + healthy_endpoints, + unhealthy_endpoints, + ) = await _run_direct_health_check_with_instrumentation( _llm_model_list, health_check_details, health_check_concurrency, instrumentation_context, ) + else: + ( + healthy_endpoints, + unhealthy_endpoints, + ) = await _run_direct_health_check_with_instrumentation( + _llm_model_list, + health_check_details, + health_check_concurrency, + instrumentation_context, ) # Update the global variable with the health check results diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/proxy_unit_tests/test_jwt_key_mapping.py index e44365897ad..7d3e7371b17 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/proxy_unit_tests/test_jwt_key_mapping.py @@ -2,18 +2,18 @@ import pytest import sys import os from unittest.mock import AsyncMock, MagicMock, patch -from fastapi import Request -from starlette.datastructures import URL -import litellm # Add project root to sys.path sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../.."))) -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth, _resolve_jwt_to_virtual_key -from litellm.proxy.auth.handle_jwt import JWTHandler, JWTAuthManager +from litellm.proxy.auth.user_api_key_auth import ( + _resolve_jwt_to_virtual_key, +) +from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy._types import LiteLLM_JWTAuth, UserAPIKeyAuth from litellm.caching.caching import DualCache + @pytest.mark.asyncio async def test_jwt_to_virtual_key_mapping_resolution(): """ @@ -21,42 +21,43 @@ async def test_jwt_to_virtual_key_mapping_resolution(): """ jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - virtual_key_claim_field="email", - virtual_key_mapping_cache_ttl=3600 + virtual_key_claim_field="email", virtual_key_mapping_cache_ttl=3600 ) - + jwt_claims = {"email": "user@example.com", "sub": "123"} - + prisma_client = MagicMock() prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock() - + # Mock finding a mapping mock_mapping = MagicMock() mock_mapping.token = "sk-1234" mock_mapping.is_active = True prisma_client.db.litellm_jwtkeymapping.find_first.return_value = mock_mapping - + # Mock getting the key object mock_key_obj = UserAPIKeyAuth(token="sk-1234", team_id="team1") - + user_api_key_cache = DualCache() - + # Use patch to mock get_key_object in the module where it's used - with patch("litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock) as mock_get_key: + with patch( + "litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock + ) as mock_get_key: mock_get_key.return_value = mock_key_obj - + result = await _resolve_jwt_to_virtual_key( jwt_claims=jwt_claims, jwt_handler=jwt_handler, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=None, - proxy_logging_obj=None + proxy_logging_obj=None, ) - + assert result == mock_key_obj prisma_client.db.litellm_jwtkeymapping.find_first.assert_called_once() - + # Test Cache hit prisma_client.db.litellm_jwtkeymapping.find_first.reset_mock() result_cached = await _resolve_jwt_to_virtual_key( @@ -65,11 +66,12 @@ async def test_jwt_to_virtual_key_mapping_resolution(): prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=None, - proxy_logging_obj=None + proxy_logging_obj=None, ) assert result_cached == mock_key_obj prisma_client.db.litellm_jwtkeymapping.find_first.assert_not_called() + @pytest.mark.asyncio async def test_jwt_to_virtual_key_mapping_no_mapping(): """ @@ -78,26 +80,28 @@ async def test_jwt_to_virtual_key_mapping_no_mapping(): jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(virtual_key_claim_field="email") jwt_claims = {"email": "unknown@example.com"} - + prisma_client = MagicMock() prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock() prisma_client.db.litellm_jwtkeymapping.find_first.return_value = None - + # Mock get_key_object just in case - with patch("litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock) as mock_get_key: + with patch( + "litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock + ): user_api_key_cache = DualCache() - + result = await _resolve_jwt_to_virtual_key( jwt_claims=jwt_claims, jwt_handler=jwt_handler, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=None, - proxy_logging_obj=None + proxy_logging_obj=None, ) - + assert result is None - + # Test Negative Cache hit prisma_client.db.litellm_jwtkeymapping.find_first.reset_mock() result_cached = await _resolve_jwt_to_virtual_key( @@ -106,7 +110,7 @@ async def test_jwt_to_virtual_key_mapping_no_mapping(): prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=None, - proxy_logging_obj=None + proxy_logging_obj=None, ) assert result_cached is None prisma_client.db.litellm_jwtkeymapping.find_first.assert_not_called()