feat reaq changes

This commit is contained in:
Harshit28j 2026-02-28 13:40:53 +05:30
parent 9dc085694c
commit 465adce872
5 changed files with 119 additions and 93 deletions

View file

@ -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

View file

@ -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,

View file

@ -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")

View file

@ -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

View file

@ -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()