mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat reaq changes
This commit is contained in:
parent
9dc085694c
commit
465adce872
5 changed files with 119 additions and 93 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue