From de23b3ec1d59455eb25477f20511928b9c03247f Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 31 Mar 2026 15:19:38 -0700 Subject: [PATCH] =?UTF-8?q?feat(proxy):=20JWT-to-virtual-key=20mapping=20i?= =?UTF-8?q?mprovements=20(P0=E2=80=93P2)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Gap #1: block key/update on jwt-bound keys for non-admins (403) - Gap #2: stamp jwt_bound metadata + restrict allowed_routes on mapping creation - Gap #3: /jwt_client/new unified endpoint — atomically creates key + mapping, returns cleartext key - Gap #4: unregistered_jwt_client_behavior config (reject/fallback_team_mapping/auto_register) - Gap #5: issuer column on LiteLLM_JWTKeyMapping for multi-IdP support - Gap #6: /jwt/key/mapping/info and /jwt_client/update expose virtual key fields - 31 unit tests covering all gaps --- litellm/proxy/_types.py | 156 +++-- litellm/proxy/auth/auth_checks.py | 11 +- litellm/proxy/auth/route_checks.py | 8 +- litellm/proxy/auth/user_api_key_auth.py | 105 +++- .../jwt_key_mapping_endpoints.py | 246 +++++++- .../key_management_endpoints.py | 62 +- litellm/proxy/schema.prisma | 5 +- .../proxy_unit_tests/test_jwt_key_mapping.py | 547 +++++++++++++++++- 8 files changed, 1050 insertions(+), 90 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8faf36df4c6..2252bbe082b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -884,9 +884,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): allowed_cache_controls: Optional[list] = [] config: Optional[dict] = {} permissions: Optional[dict] = {} - model_max_budget: Optional[ - dict - ] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} + model_max_budget: Optional[dict] = ( + {} + ) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} model_config = ConfigDict(protected_namespaces=()) model_rpm_limit: Optional[dict] = None @@ -938,6 +938,7 @@ class LiteLLMKeyType(str, enum.Enum): MANAGEMENT = "management" # Can call management routes (user/team/key management) READ_ONLY = "read_only" # Can only call info/read routes DEFAULT = "default" # Uses default allowed routes + JWT_CLIENT = "jwt_client" # JWT-mapped service identity — LLM API routes only, no key management class GenerateKeyRequest(KeyRequestBase): @@ -1028,9 +1029,9 @@ class RegenerateKeyRequest(GenerateKeyRequest): spend: Optional[float] = None metadata: Optional[dict] = None new_master_key: Optional[str] = None - grace_period: Optional[ - str - ] = None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke + grace_period: Optional[str] = ( + None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke + ) class ResetSpendRequest(LiteLLMPydanticObjectBase): @@ -1540,12 +1541,12 @@ class NewCustomerRequest(BudgetNewRequest): blocked: bool = False # allow/disallow requests for this end-user budget_id: Optional[str] = None # give either a budget_id or max_budget spend: Optional[float] = None - allowed_model_region: Optional[ - AllowedModelRegion - ] = None # require all user requests to use models in this specific region - default_model: Optional[ - str - ] = None # if no equivalent model in allowed region - default all requests to this model + allowed_model_region: Optional[AllowedModelRegion] = ( + None # require all user requests to use models in this specific region + ) + default_model: Optional[str] = ( + None # if no equivalent model in allowed region - default all requests to this model + ) object_permission: Optional[LiteLLM_ObjectPermissionBase] = None @model_validator(mode="before") @@ -1568,12 +1569,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase): blocked: bool = False # allow/disallow requests for this end-user max_budget: Optional[float] = None budget_id: Optional[str] = None # give either a budget_id or max_budget - allowed_model_region: Optional[ - AllowedModelRegion - ] = None # require all user requests to use models in this specific region - default_model: Optional[ - str - ] = None # if no equivalent model in allowed region - default all requests to this model + allowed_model_region: Optional[AllowedModelRegion] = ( + None # require all user requests to use models in this specific region + ) + default_model: Optional[str] = ( + None # if no equivalent model in allowed region - default all requests to this model + ) object_permission: Optional[LiteLLM_ObjectPermissionBase] = None @@ -1663,15 +1664,15 @@ class NewTeamRequest(TeamBase): ] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm model_tpm_limit: Optional[Dict[str, int]] = None - team_member_budget: Optional[ - float - ] = None # allow user to set a budget for all team members - team_member_rpm_limit: Optional[ - int - ] = None # allow user to set RPM limit for all team members - team_member_tpm_limit: Optional[ - int - ] = None # allow user to set TPM limit for all team members + team_member_budget: Optional[float] = ( + None # allow user to set a budget for all team members + ) + team_member_rpm_limit: Optional[int] = ( + None # allow user to set RPM limit for all team members + ) + team_member_tpm_limit: Optional[int] = ( + None # allow user to set TPM limit for all team members + ) team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m" team_member_budget_duration: Optional[str] = None # e.g. "30d", "1mo" allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None @@ -1768,9 +1769,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase): class AddTeamCallback(LiteLLMPydanticObjectBase): callback_name: str - callback_type: Optional[ - Literal["success", "failure", "success_and_failure"] - ] = "success_and_failure" + callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = ( + "success_and_failure" + ) callback_vars: Dict[str, str] @model_validator(mode="before") @@ -2110,9 +2111,9 @@ class ConfigList(LiteLLMPydanticObjectBase): stored_in_db: Optional[bool] field_default_value: Any premium_field: bool = False - nested_fields: Optional[ - List[FieldDetail] - ] = None # For nested dictionary or Pydantic fields + nested_fields: Optional[List[FieldDetail]] = ( + None # For nested dictionary or Pydantic fields + ) class UserHeaderMapping(LiteLLMPydanticObjectBase): @@ -2470,9 +2471,9 @@ class UserAPIKeyAuth( user_max_budget: Optional[float] = None request_route: Optional[str] = None user: Optional[Any] = None # Expanded user object when expand=user is used - created_by_user: Optional[ - Any - ] = None # Expanded created_by user when expand=user is used + created_by_user: Optional[Any] = ( + None # Expanded created_by user when expand=user is used + ) end_user_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None # Decoded upstream IdP claims (groups, roles, etc.) propagated by JWT auth machinery # and forwarded into outbound tokens by guardrails such as MCPJWTSigner. @@ -2611,9 +2612,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase): budget_id: Optional[str] = None created_at: datetime updated_at: datetime - user: Optional[ - Any - ] = None # You might want to replace 'Any' with a more specific type if available + user: Optional[Any] = ( + None # You might want to replace 'Any' with a more specific type if available + ) litellm_budget_table: Optional[LiteLLM_BudgetTable] = None user_email: Optional[str] = None @@ -3764,9 +3765,9 @@ class TeamModelDeleteRequest(BaseModel): # Organization Member Requests class OrganizationMemberAddRequest(OrgMemberAddRequest): organization_id: str - max_budget_in_organization: Optional[ - float - ] = None # Users max budget within the organization + max_budget_in_organization: Optional[float] = ( + None # Users max budget within the organization + ) class OrganizationMemberDeleteRequest(MemberDeleteRequest): @@ -3846,6 +3847,7 @@ class KeyHealthResponse(TypedDict, total=False): class CreateJWTKeyMappingRequest(LiteLLMPydanticObjectBase): jwt_claim_name: str jwt_claim_value: str + issuer: str = "" key: str description: Optional[str] = None @@ -3865,12 +3867,55 @@ class JWTKeyMappingResponse(LiteLLMPydanticObjectBase): id: str jwt_claim_name: str jwt_claim_value: str + issuer: str = "" description: Optional[str] = None is_active: bool created_at: datetime updated_at: datetime created_by: Optional[str] = None updated_by: Optional[str] = None + # Virtual key fields — populated by info/unified endpoints + models: Optional[list] = None + max_budget: Optional[float] = None + budget_duration: Optional[str] = None + tpm_limit: Optional[int] = None + rpm_limit: Optional[int] = None + team_id: Optional[str] = None + key_alias: Optional[str] = None + spend: Optional[float] = None + expires: Optional[datetime] = None + # Cleartext key — only populated on creation (never returned again) + key: Optional[str] = None + + +class CreateJWTClientRequest(LiteLLMPydanticObjectBase): + """Single-call request to create a virtual key + JWT mapping atomically.""" + + jwt_claim_name: str + jwt_claim_value: str + issuer: str = "" + description: Optional[str] = None + # Virtual key configuration + models: Optional[list] = [] + max_budget: Optional[float] = None + budget_duration: Optional[str] = None + tpm_limit: Optional[int] = None + rpm_limit: Optional[int] = None + team_id: Optional[str] = None + key_alias: Optional[str] = None + duration: Optional[str] = None + metadata: Optional[dict] = {} + + +class JWTClientAutoRegisterDefaults(LiteLLMPydanticObjectBase): + """Default virtual key settings applied when auto-registering unknown JWT clients.""" + + models: Optional[list] = None + max_budget: Optional[float] = None + budget_duration: Optional[str] = None + tpm_limit: Optional[int] = None + rpm_limit: Optional[int] = None + team_id: Optional[str] = None class SpecialHeaders(enum.Enum): @@ -4017,9 +4062,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase): Maps provider names to their budget configs. """ - providers: Dict[ - str, ProviderBudgetResponseObject - ] = {} # Dictionary mapping provider names to their budget configurations + providers: Dict[str, ProviderBudgetResponseObject] = ( + {} + ) # Dictionary mapping provider names to their budget configurations class ProxyStateVariables(TypedDict): @@ -4163,9 +4208,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): enforce_rbac: bool = False roles_jwt_field: Optional[str] = None # v2 on role mappings role_mappings: Optional[List[RoleMapping]] = None - object_id_jwt_field: Optional[ - str - ] = None # can be either user / team, inferred from the role mapping + object_id_jwt_field: Optional[str] = ( + None # can be either user / team, inferred from the role mapping + ) scope_mappings: Optional[List[ScopeMapping]] = None enforce_scope_based_access: bool = False enforce_team_based_model_access: bool = False @@ -4198,6 +4243,21 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): default=300, description="TTL (seconds) for caching JWT-to-virtual-key mapping lookups.", ) + unregistered_jwt_client_behavior: Literal[ + "reject", "fallback_team_mapping", "auto_register" + ] = Field( + default="fallback_team_mapping", + description=( + "Behavior when a JWT arrives with no virtual-key mapping. " + "'reject' → 403. " + "'fallback_team_mapping' → standard JWT team auth (default). " + "'auto_register' → create a mapping + virtual key on first encounter." + ), + ) + auto_register_defaults: Optional["JWTClientAutoRegisterDefaults"] = Field( + default=None, + description="Default virtual key settings used when auto_register creates keys for unknown JWT clients.", + ) ######################################################### def __init__(self, **kwargs: Any) -> None: diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index c00a351bdd8..f2ad3308d6c 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -8,6 +8,7 @@ Run checks for: 2. If user is in budget 3. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget """ + import asyncio import re import time @@ -414,9 +415,9 @@ async def common_checks( # noqa: PLR0915 model=_model, team_object=team_object, llm_router=llm_router, - team_model_aliases=valid_token.team_model_aliases - if valid_token - else None, + team_model_aliases=( + valid_token.team_model_aliases if valid_token else None + ), ): raise ProxyException( message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}", @@ -2146,16 +2147,20 @@ async def get_jwt_key_mapping_object( jwt_claim_name: str, jwt_claim_value: str, prisma_client: PrismaClient, + issuer: str = "", ) -> Optional[str]: """ Lookup a JWT-to-virtual-key mapping from the database. Returns the hashed token (str) if a matching active mapping is found, else None. + The `issuer` parameter (JWT "iss" claim) is used to disambiguate identical + claim values across different identity providers. """ mapping = await prisma_client.db.litellm_jwtkeymapping.find_first( where={ "jwt_claim_name": jwt_claim_name, "jwt_claim_value": jwt_claim_value, + "issuer": issuer, "is_active": True, } ) diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 26bbdef3090..75818c214dc 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -209,7 +209,13 @@ class RouteChecks: route=route, allowed_routes=LiteLLMRoutes.internal_user_routes.value ) ): - pass + # JWT-bound sessions (resolved from a JWT-to-virtual-key mapping) are + # restricted to LLM API routes only. Non-LLM routes reach this branch + # only because they passed the is_llm_api_route check above, so deny here. + if valid_token.jwt_claims is not None: + RouteChecks._raise_admin_only_route_exception( + user_obj=user_obj, route=route + ) elif _user_is_org_admin( request_data=request_data, user_object=user_obj ) and RouteChecks.check_route_access( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 53ae08aefb1..28cec61de16 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -471,7 +471,8 @@ async def _resolve_jwt_to_virtual_key( ) return None - cache_key = f"jwt_key_mapping:{virtual_key_claim_field}:{claim_value}" + issuer = str(jwt_claims.get("iss", "")) + cache_key = f"jwt_key_mapping:{virtual_key_claim_field}:{claim_value}:{issuer}" cached_mapping = await user_api_key_cache.async_get_cache(cache_key) if cached_mapping == "__NO_MAPPING__": @@ -492,6 +493,7 @@ async def _resolve_jwt_to_virtual_key( jwt_claim_name=virtual_key_claim_field, jwt_claim_value=str(claim_value), prisma_client=prisma_client, + issuer=issuer, ) if token_hash is not None: @@ -507,7 +509,26 @@ async def _resolve_jwt_to_virtual_key( parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) + + behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior + if behavior == "reject": + raise HTTPException( + status_code=403, + detail=f"JWT client not registered. No mapping found for claim '{virtual_key_claim_field}' = '{claim_value}'.", + ) + elif behavior == "auto_register": + return await _auto_register_jwt_client( + jwt_claim_name=virtual_key_claim_field, + jwt_claim_value=str(claim_value), + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + cache_key=cache_key, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) else: + # "fallback_team_mapping" — default: fall through to standard JWT auth await user_api_key_cache.async_set_cache( key=cache_key, value="__NO_MAPPING__", @@ -516,6 +537,76 @@ async def _resolve_jwt_to_virtual_key( return None +async def _auto_register_jwt_client( + jwt_claim_name: str, + jwt_claim_value: str, + jwt_handler: "JWTHandler", + prisma_client: "PrismaClient", + user_api_key_cache: DualCache, + cache_key: str, + parent_otel_span: Optional["Span"], + proxy_logging_obj: "ProxyLogging", +) -> Optional["UserAPIKeyAuth"]: + """ + Auto-register a new JWT client by creating a virtual key + mapping on first encounter. + Called when unregistered_jwt_client_behavior == 'auto_register'. + """ + from litellm.proxy._types import hash_token + from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_key_helper_fn, + ) + + defaults = jwt_handler.litellm_jwtauth.auto_register_defaults + jwt_metadata = { + "jwt_bound": True, + "jwt_claim_name": jwt_claim_name, + "jwt_claim_value": jwt_claim_value, + } + key_data = await generate_key_helper_fn( + request_type="key", + models=list(defaults.models or []) if defaults else [], + metadata=jwt_metadata, + key_max_budget=defaults.max_budget if defaults else None, + key_budget_duration=defaults.budget_duration if defaults else None, + tpm_limit=defaults.tpm_limit if defaults else None, + rpm_limit=defaults.rpm_limit if defaults else None, + team_id=defaults.team_id if defaults else None, + allowed_routes=["llm_api_routes"], + ) + + cleartext_token: str = key_data["token"] + hashed_key = hash_token(cleartext_token) + + issuer = "" # populated by caller if available; extend as needed + try: + await prisma_client.db.litellm_jwtkeymapping.create( + data={ + "jwt_claim_name": jwt_claim_name, + "jwt_claim_value": jwt_claim_value, + "issuer": issuer, + "token": hashed_key, + } + ) + except Exception: + verbose_proxy_logger.warning( + f"JWT auto-register: failed to create mapping for {jwt_claim_name}={jwt_claim_value}. " + "Key was created but mapping row could not be saved." + ) + + await user_api_key_cache.async_set_cache( + key=cache_key, + value=hashed_key, + ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl, + ) + return await get_key_object( + hashed_token=hashed_key, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + + async def _user_api_key_auth_builder( # noqa: PLR0915 request: Request, api_key: str, @@ -924,9 +1015,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 route=route, ) if _end_user_object is not None: - end_user_params[ - "allowed_model_region" - ] = _end_user_object.allowed_model_region + end_user_params["allowed_model_region"] = ( + _end_user_object.allowed_model_region + ) if _end_user_object.litellm_budget_table is not None: _apply_budget_limits_to_end_user_params( end_user_params=end_user_params, @@ -1516,9 +1607,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if _end_user_object is not None: valid_token_dict.update(end_user_params) - valid_token_dict[ - "end_user_object_permission" - ] = _end_user_object.object_permission + valid_token_dict["end_user_object_permission"] = ( + _end_user_object.object_permission + ) # check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions # sso/login, ui/login, /key functions and /user functions diff --git a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py index e474cb7d155..7d5508fc3d4 100644 --- a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py +++ b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py @@ -1,6 +1,10 @@ +import json +from typing import Optional + from fastapi import APIRouter, Depends, HTTPException, Query from litellm.proxy._types import ( + CreateJWTClientRequest, CreateJWTKeyMappingRequest, DeleteJWTKeyMappingRequest, JWTKeyMappingResponse, @@ -14,12 +18,17 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth router = APIRouter() -def _to_response(mapping) -> JWTKeyMappingResponse: - """Convert a Prisma mapping object to a safe response (no hashed token).""" - return JWTKeyMappingResponse( +def _to_response(mapping, key_row=None) -> JWTKeyMappingResponse: + """Convert a Prisma mapping object to a safe response (no hashed token). + + Optionally accepts the linked LiteLLM_VerificationToken row to populate + virtual key fields (models, budget, rate limits, etc.). + """ + resp = JWTKeyMappingResponse( id=mapping.id, jwt_claim_name=mapping.jwt_claim_name, jwt_claim_value=mapping.jwt_claim_value, + issuer=getattr(mapping, "issuer", ""), description=mapping.description, is_active=mapping.is_active, created_at=mapping.created_at, @@ -27,6 +36,17 @@ def _to_response(mapping) -> JWTKeyMappingResponse: created_by=mapping.created_by, updated_by=mapping.updated_by, ) + if key_row is not None: + resp.models = key_row.models or [] + resp.max_budget = key_row.max_budget + resp.budget_duration = key_row.budget_duration + resp.tpm_limit = key_row.tpm_limit + resp.rpm_limit = key_row.rpm_limit + resp.team_id = key_row.team_id + resp.key_alias = key_row.key_alias + resp.spend = key_row.spend + resp.expires = key_row.expires + return resp @router.post( @@ -53,6 +73,7 @@ async def create_jwt_key_mapping( create_data = { "jwt_claim_name": data.jwt_claim_name, "jwt_claim_value": data.jwt_claim_value, + "issuer": data.issuer, "token": hashed_key, "created_by": user_api_key_dict.user_id, "updated_by": user_api_key_dict.user_id, @@ -64,8 +85,34 @@ async def create_jwt_key_mapping( data=create_data ) - # Invalidate cache - cache_key = f"jwt_key_mapping:{data.jwt_claim_name}:{data.jwt_claim_value}" + # Stamp jwt_bound metadata and restrict routes on the mapped key + existing_key = await prisma_client.db.litellm_verificationtoken.find_first( + where={"token": hashed_key} + ) + if existing_key is not None: + existing_metadata = {} + if existing_key.metadata: + raw = existing_key.metadata + if isinstance(raw, str): + existing_metadata = json.loads(raw) + elif isinstance(raw, dict): + existing_metadata = raw + jwt_metadata = { + **existing_metadata, + "jwt_bound": True, + "jwt_claim_name": data.jwt_claim_name, + "jwt_claim_value": data.jwt_claim_value, + } + await prisma_client.db.litellm_verificationtoken.update( + where={"token": hashed_key}, + data={ + "metadata": json.dumps(jwt_metadata), + "allowed_routes": ["llm_api_routes"], + }, + ) + + # Invalidate cache (include issuer in cache key) + cache_key = f"jwt_key_mapping:{data.jwt_claim_name}:{data.jwt_claim_value}:{data.issuer}" await user_api_key_cache.async_delete_cache(cache_key) return _to_response(new_mapping) @@ -119,7 +166,8 @@ async def update_jwt_key_mapping( if old_mapping is None: raise HTTPException(status_code=404, detail="Mapping not found") - cache_key = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}" + old_issuer = getattr(old_mapping, "issuer", "") + cache_key = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}:{old_issuer}" await user_api_key_cache.async_delete_cache(cache_key) updated_mapping = await prisma_client.db.litellm_jwtkeymapping.update( @@ -127,7 +175,8 @@ async def update_jwt_key_mapping( ) # Invalidate new cache key if claim fields changed - cache_key = f"jwt_key_mapping:{updated_mapping.jwt_claim_name}:{updated_mapping.jwt_claim_value}" + new_issuer = getattr(updated_mapping, "issuer", "") + cache_key = f"jwt_key_mapping:{updated_mapping.jwt_claim_name}:{updated_mapping.jwt_claim_value}:{new_issuer}" await user_api_key_cache.async_delete_cache(cache_key) return _to_response(updated_mapping) @@ -172,7 +221,8 @@ async def delete_jwt_key_mapping( if old_mapping is None: raise HTTPException(status_code=404, detail="Mapping not found") - cache_key = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}" + del_issuer = getattr(old_mapping, "issuer", "") + cache_key = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}:{del_issuer}" await user_api_key_cache.async_delete_cache(cache_key) await prisma_client.db.litellm_jwtkeymapping.delete(where={"id": data.id}) @@ -247,10 +297,188 @@ async def info_jwt_key_mapping( ) if mapping is None: raise HTTPException(status_code=404, detail="Mapping not found") - return _to_response(mapping) + key_row = await prisma_client.db.litellm_verificationtoken.find_first( + where={"token": mapping.token} + ) + return _to_response(mapping, key_row=key_row) except HTTPException: raise except Exception: raise HTTPException( status_code=500, detail="Failed to get JWT key mapping info." ) + + +@router.post( + "/jwt_client/new", + tags=["JWT Key Mapping"], + response_model=JWTKeyMappingResponse, +) +async def create_jwt_client( + data: CreateJWTClientRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Atomically create a virtual key + JWT mapping in one call. + The key is automatically typed as JWT_CLIENT (LLM API routes only) + and stamped with jwt_bound metadata. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_key_helper_fn, + ) + 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 clients" + ) + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + jwt_metadata = { + **(data.metadata or {}), + "jwt_bound": True, + "jwt_claim_name": data.jwt_claim_name, + "jwt_claim_value": data.jwt_claim_value, + } + + key_data = await generate_key_helper_fn( + request_type="key", + table_name="key", # skip user-table upsert (user_id may be None for master key) + models=data.models or [], + metadata=jwt_metadata, + key_max_budget=data.max_budget, + key_budget_duration=data.budget_duration, + tpm_limit=data.tpm_limit, + rpm_limit=data.rpm_limit, + team_id=data.team_id, + key_alias=data.key_alias, + duration=data.duration, + allowed_routes=["llm_api_routes"], + created_by=user_api_key_dict.user_id, + updated_by=user_api_key_dict.user_id, + ) + + cleartext_token: str = key_data["token"] + hashed_key = hash_token(cleartext_token) + + try: + create_data = { + "jwt_claim_name": data.jwt_claim_name, + "jwt_claim_value": data.jwt_claim_value, + "issuer": data.issuer, + "token": hashed_key, + "created_by": user_api_key_dict.user_id, + "updated_by": user_api_key_dict.user_id, + } + if data.description is not None: + create_data["description"] = data.description + + new_mapping = await prisma_client.db.litellm_jwtkeymapping.create( + data=create_data + ) + except Exception as e: + # Best-effort cleanup of the orphaned key + try: + await prisma_client.db.litellm_verificationtoken.delete( + where={"token": hashed_key} + ) + except Exception: + pass + error_str = str(e).lower() + if "unique" in error_str or "p2002" in error_str: + raise HTTPException( + status_code=409, + detail=f"A JWT client for claim '{data.jwt_claim_name}' = '{data.jwt_claim_value}' already exists.", + ) + raise HTTPException(status_code=500, detail="Failed to create JWT client.") + + cache_key = ( + f"jwt_key_mapping:{data.jwt_claim_name}:{data.jwt_claim_value}:{data.issuer}" + ) + await user_api_key_cache.async_delete_cache(cache_key) + + key_row = await prisma_client.db.litellm_verificationtoken.find_first( + where={"token": hashed_key} + ) + resp = _to_response(new_mapping, key_row=key_row) + resp.key = cleartext_token # return cleartext key only on creation + return resp + + +@router.post( + "/jwt_client/update", + tags=["JWT Key Mapping"], + response_model=JWTKeyMappingResponse, +) +async def update_jwt_client( + id: str, + models: Optional[list] = None, + max_budget: Optional[float] = None, + budget_duration: Optional[str] = None, + tpm_limit: Optional[int] = None, + rpm_limit: Optional[int] = None, + description: Optional[str] = None, + is_active: Optional[bool] = None, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Update a JWT client's virtual key configuration and/or mapping metadata. + """ + 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 clients" + ) + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + mapping = await prisma_client.db.litellm_jwtkeymapping.find_unique(where={"id": id}) + if mapping is None: + raise HTTPException(status_code=404, detail="JWT client not found") + + # Update mapping metadata + mapping_update: dict = {"updated_by": user_api_key_dict.user_id} + if description is not None: + mapping_update["description"] = description + if is_active is not None: + mapping_update["is_active"] = is_active + + updated_mapping = await prisma_client.db.litellm_jwtkeymapping.update( + where={"id": id}, data=mapping_update + ) + + # Update underlying virtual key + key_update: dict = {} + if models is not None: + key_update["models"] = models + if max_budget is not None: + key_update["max_budget"] = max_budget + if budget_duration is not None: + key_update["budget_duration"] = budget_duration + if tpm_limit is not None: + key_update["tpm_limit"] = tpm_limit + if rpm_limit is not None: + key_update["rpm_limit"] = rpm_limit + + key_row = None + if key_update: + key_row = await prisma_client.db.litellm_verificationtoken.update( + where={"token": mapping.token}, + data=key_update, + ) + else: + key_row = await prisma_client.db.litellm_verificationtoken.find_first( + where={"token": mapping.token} + ) + + issuer = getattr(mapping, "issuer", "") + cache_key = ( + f"jwt_key_mapping:{mapping.jwt_claim_name}:{mapping.jwt_claim_value}:{issuer}" + ) + await user_api_key_cache.async_delete_cache(cache_key) + + return _to_response(updated_mapping, key_row=key_row) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 831922ec3f9..180c50206ed 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -453,6 +453,8 @@ def handle_key_type(data: GenerateKeyRequest, data_json: dict) -> dict: data_json["allowed_routes"] = ["management_routes"] elif key_type == LiteLLMKeyType.READ_ONLY: data_json["allowed_routes"] = ["info_routes"] + elif key_type == LiteLLMKeyType.JWT_CLIENT: + data_json["allowed_routes"] = ["llm_api_routes"] return data_json @@ -732,9 +734,9 @@ async def _common_key_generation_helper( # noqa: PLR0915 request_type="key", **data_json, table_name="key" ) - response[ - "soft_budget" - ] = data.soft_budget # include the user-input soft budget in the response + response["soft_budget"] = ( + data.soft_budget + ) # include the user-input soft budget in the response response = GenerateKeyResponse(**response) @@ -2116,6 +2118,19 @@ async def update_key_fn( prisma_client=prisma_client, ) + # JWT-bound keys can only be modified by proxy admins + _key_metadata = existing_key_row.metadata or {} + if isinstance(_key_metadata, str): + _key_metadata = json.loads(_key_metadata) + if ( + _key_metadata.get("jwt_bound") + and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + ): + raise HTTPException( + status_code=403, + detail="JWT-bound keys can only be modified by proxy admins", + ) + await _validate_update_key_data( data=data, existing_key_row=existing_key_row, @@ -3103,6 +3118,15 @@ async def can_modify_verification_token( if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return True + # 1b. JWT-bound keys cannot be modified by non-admin sessions + key_metadata = key_info.metadata or {} + if isinstance(key_metadata, str): + import json as _json + + key_metadata = _json.loads(key_metadata) + if key_metadata.get("jwt_bound"): + return False + # 2. Internal jobs service account can modify any key (for auto-rotation) if user_api_key_dict.api_key == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME: return True @@ -3175,10 +3199,10 @@ async def delete_verification_tokens( try: if prisma_client: tokens = [_hash_token_if_needed(token=key) for key in tokens] - _keys_being_deleted: List[ - LiteLLM_VerificationToken - ] = await prisma_client.db.litellm_verificationtoken.find_many( - where={"token": {"in": tokens}} + _keys_being_deleted: List[LiteLLM_VerificationToken] = ( + await prisma_client.db.litellm_verificationtoken.find_many( + where={"token": {"in": tokens}} + ) ) if len(_keys_being_deleted) == 0: @@ -3378,9 +3402,9 @@ async def _rotate_master_key( # noqa: PLR0915 from litellm.proxy.proxy_server import proxy_config try: - models: Optional[ - List - ] = await prisma_client.db.litellm_proxymodeltable.find_many() + models: Optional[List] = ( + await prisma_client.db.litellm_proxymodeltable.find_many() + ) except Exception: models = None # 2. process model table @@ -4020,11 +4044,11 @@ async def validate_key_list_check( param="user_id", code=status.HTTP_403_FORBIDDEN, ) - complete_user_info_db_obj: Optional[ - BaseModel - ] = await prisma_client.db.litellm_usertable.find_unique( - where={"user_id": user_api_key_dict.user_id}, - include={"organization_memberships": True}, + complete_user_info_db_obj: Optional[BaseModel] = ( + await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": user_api_key_dict.user_id}, + include={"organization_memberships": True}, + ) ) if complete_user_info_db_obj is None: @@ -4107,10 +4131,10 @@ async def _fetch_user_team_objects( if complete_user_info is None or not complete_user_info.teams: return [] - teams: Optional[ - List[BaseModel] - ] = await prisma_client.db.litellm_teamtable.find_many( - where={"team_id": {"in": complete_user_info.teams}} + teams: Optional[List[BaseModel]] = ( + await prisma_client.db.litellm_teamtable.find_many( + where={"team_id": {"in": complete_user_info.teams}} + ) ) if teams is None: return [] diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 46be6b31e1f..519243e166d 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -408,6 +408,7 @@ model LiteLLM_JWTKeyMapping { id String @id @default(uuid()) jwt_claim_name String // e.g. "sub", "email" jwt_claim_value String // The claim value to match + issuer String @default("") // JWT "iss" claim — differentiates identical claims across IdPs token String // Hashed virtual key (FK) description String? is_active Boolean @default(true) @@ -418,8 +419,8 @@ model LiteLLM_JWTKeyMapping { litellm_verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token]) - @@unique([jwt_claim_name, jwt_claim_value]) - @@index([jwt_claim_name, jwt_claim_value, is_active]) + @@unique([jwt_claim_name, jwt_claim_value, issuer]) + @@index([jwt_claim_name, jwt_claim_value, issuer, is_active]) } // Deprecated keys during grace period - allows old key to work until revoke_at diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/proxy_unit_tests/test_jwt_key_mapping.py index 66e5b3839bc..55903231a95 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/proxy_unit_tests/test_jwt_key_mapping.py @@ -222,6 +222,7 @@ def test_to_response_excludes_token(): mock_mapping.id = "mapping-1" mock_mapping.jwt_claim_name = "email" mock_mapping.jwt_claim_value = "user@example.com" + mock_mapping.issuer = "" mock_mapping.token = "hashed_secret_value" mock_mapping.description = "test" mock_mapping.is_active = True @@ -265,6 +266,10 @@ def _mock_prisma(): prisma.db.litellm_jwtkeymapping.update = AsyncMock() prisma.db.litellm_jwtkeymapping.delete = AsyncMock() prisma.db.litellm_jwtkeymapping.count = AsyncMock(return_value=0) + # Needed by create_jwt_key_mapping (stamps jwt_bound metadata) and info endpoint + prisma.db.litellm_verificationtoken.find_first = AsyncMock(return_value=None) + prisma.db.litellm_verificationtoken.update = AsyncMock() + prisma.db.litellm_verificationtoken.delete = AsyncMock() return prisma @@ -272,12 +277,14 @@ def _mock_mapping( id="mapping-1", claim_name="email", claim_value="user@example.com", + issuer="", ): now = datetime.now(timezone.utc) m = MagicMock() m.id = id m.jwt_claim_name = claim_name m.jwt_claim_value = claim_value + m.issuer = issuer m.token = "hashed_token" m.description = None m.is_active = True @@ -398,8 +405,11 @@ async def test_info_returns_404_when_not_found(): """Getting info for non-existent mapping should return 404.""" mock_prisma = _mock_prisma() mock_prisma.db.litellm_jwtkeymapping.find_unique.return_value = None + mock_cache = AsyncMock() - with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_cache + ): with pytest.raises(HTTPException) as exc_info: await info_jwt_key_mapping(id="nonexistent-id", user_api_key_dict=_make_admin_auth()) assert exc_info.value.status_code == 404 @@ -425,3 +435,538 @@ async def test_create_success_returns_response_without_token(): assert isinstance(result, JWTKeyMappingResponse) assert "token" not in result.model_fields assert result.jwt_claim_name == "email" + + +# ────────────────────────────────────────────── +# Block 1: Gap #2 — JWT-bound metadata stamping +# ────────────────────────────────────────────── + + +def test_handle_key_type_jwt_client(): + """JWT_CLIENT key type should resolve to llm_api_routes.""" + from litellm.proxy._types import LiteLLMKeyType + from litellm.proxy.management_endpoints.key_management_endpoints import ( + handle_key_type, + ) + from litellm.proxy._types import GenerateKeyRequest + + data = GenerateKeyRequest(key_type=LiteLLMKeyType.JWT_CLIENT) + data_json = data.model_dump() + result = handle_key_type(data=data, data_json=data_json) + assert result["allowed_routes"] == ["llm_api_routes"] + + +@pytest.mark.asyncio +async def test_create_mapping_stamps_jwt_bound_metadata(): + """create_jwt_key_mapping should update the key's metadata with jwt_bound=True.""" + from litellm.proxy._types import CreateJWTKeyMappingRequest + import json as _json + + mock_prisma = _mock_prisma() + mock_mapping = _mock_mapping() + mock_prisma.db.litellm_jwtkeymapping.create.return_value = mock_mapping + + # Simulate key with no existing metadata + mock_key_row = MagicMock() + mock_key_row.metadata = "{}" + mock_prisma.db.litellm_verificationtoken.find_first.return_value = mock_key_row + mock_cache = AsyncMock() + + data = CreateJWTKeyMappingRequest( + jwt_claim_name="sub", + jwt_claim_value="svc-a", + key="sk-test-key", + ) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_cache + ): + await create_jwt_key_mapping(data=data, user_api_key_dict=_make_admin_auth()) + + mock_prisma.db.litellm_verificationtoken.update.assert_called_once() + call_kwargs = mock_prisma.db.litellm_verificationtoken.update.call_args + update_data = call_kwargs.kwargs["data"] + stored_metadata = _json.loads(update_data["metadata"]) + assert stored_metadata["jwt_bound"] is True + assert stored_metadata["jwt_claim_name"] == "sub" + assert stored_metadata["jwt_claim_value"] == "svc-a" + assert update_data["allowed_routes"] == ["llm_api_routes"] + + +@pytest.mark.asyncio +async def test_create_mapping_preserves_existing_metadata(): + """Existing metadata keys should be preserved when stamping jwt_bound.""" + from litellm.proxy._types import CreateJWTKeyMappingRequest + import json as _json + + mock_prisma = _mock_prisma() + mock_prisma.db.litellm_jwtkeymapping.create.return_value = _mock_mapping() + + mock_key_row = MagicMock() + mock_key_row.metadata = _json.dumps({"custom_field": "keep_me"}) + mock_prisma.db.litellm_verificationtoken.find_first.return_value = mock_key_row + mock_cache = AsyncMock() + + data = CreateJWTKeyMappingRequest( + jwt_claim_name="email", + jwt_claim_value="user@example.com", + key="sk-test-key", + ) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_cache + ): + await create_jwt_key_mapping(data=data, user_api_key_dict=_make_admin_auth()) + + update_data = mock_prisma.db.litellm_verificationtoken.update.call_args.kwargs["data"] + stored_metadata = _json.loads(update_data["metadata"]) + assert stored_metadata["custom_field"] == "keep_me" + assert stored_metadata["jwt_bound"] is True + + +# ────────────────────────────────────────────── +# Block 2: Gap #1 — CRUD block on jwt-bound keys +# ────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_can_modify_jwt_bound_key_as_admin_returns_true(): + """Proxy admin can always modify JWT-bound keys.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + can_modify_verification_token, + ) + from litellm.proxy._types import LiteLLM_VerificationToken + + key_info = MagicMock(spec=LiteLLM_VerificationToken) + key_info.metadata = {"jwt_bound": True} + key_info.team_id = None + key_info.user_id = "some-user" + + admin_dict = _make_admin_auth() + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=MagicMock(), + user_api_key_dict=admin_dict, + prisma_client=MagicMock(), + ) + assert result is True + + +@pytest.mark.asyncio +async def test_can_modify_jwt_bound_key_as_internal_user_returns_false(): + """Non-admin cannot modify a JWT-bound key even if they own it.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + can_modify_verification_token, + ) + from litellm.proxy._types import LiteLLM_VerificationToken + + key_info = MagicMock(spec=LiteLLM_VerificationToken) + key_info.metadata = {"jwt_bound": True} + key_info.team_id = None + key_info.user_id = "user-123" + + non_admin = _make_non_admin_auth() + non_admin.user_id = "user-123" # same user_id as key owner + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=MagicMock(), + user_api_key_dict=non_admin, + prisma_client=MagicMock(), + ) + assert result is False + + +def test_jwt_session_blocked_from_key_management_route(): + """A JWT-authenticated session (jwt_claims set) must not reach key management routes.""" + from litellm.proxy.auth.route_checks import RouteChecks + from unittest.mock import MagicMock + + user_obj = MagicMock() + user_obj.user_id = "user-123" + + valid_token = UserAPIKeyAuth( + token="sk-bound", + user_role=LitellmUserRoles.INTERNAL_USER, + jwt_claims={"sub": "svc-a"}, # marks this as a JWT session + ) + + with pytest.raises(Exception): + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.INTERNAL_USER, + route="/key/update", + request=MagicMock(), + valid_token=valid_token, + request_data={}, + ) + + +def test_jwt_session_allowed_llm_api_route(): + """A JWT-authenticated session must be allowed to call LLM API routes.""" + from litellm.proxy.auth.route_checks import RouteChecks + + valid_token = UserAPIKeyAuth( + token="sk-bound", + user_role=LitellmUserRoles.INTERNAL_USER, + jwt_claims={"sub": "svc-a"}, + ) + + # Should not raise for an LLM API route + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=MagicMock(), + _user_role=LitellmUserRoles.INTERNAL_USER, + route="/v1/chat/completions", + request=MagicMock(), + valid_token=valid_token, + request_data={}, + ) + + +def test_non_jwt_session_internal_user_can_access_key_management(): + """Regular API key session (no jwt_claims) keeps existing internal_user access.""" + from litellm.proxy.auth.route_checks import RouteChecks + + valid_token = UserAPIKeyAuth( + token="sk-regular", + user_role=LitellmUserRoles.INTERNAL_USER, + jwt_claims=None, # not a JWT session + ) + + # Should not raise — existing behavior preserved + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=MagicMock(), + _user_role=LitellmUserRoles.INTERNAL_USER, + route="/key/update", + request=MagicMock(), + valid_token=valid_token, + request_data={}, + ) + + +# ────────────────────────────────────────────── +# Block 3: Gap #3 — Unified /jwt_client/new +# ────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_jwt_client_new_creates_key_and_mapping(): + """/jwt_client/new should call generate_key_helper_fn and create a mapping row.""" + from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import ( + create_jwt_client, + ) + from litellm.proxy._types import CreateJWTClientRequest + import json as _json + + mock_prisma = _mock_prisma() + mock_mapping = _mock_mapping(claim_name="sub", claim_value="svc-a") + mock_prisma.db.litellm_jwtkeymapping.create.return_value = mock_mapping + mock_cache = AsyncMock() + + fake_key_data = {"token": "sk-auto-generated-key"} + + data = CreateJWTClientRequest( + jwt_claim_name="sub", + jwt_claim_value="svc-a", + models=["gpt-4o"], + max_budget=10.0, + ) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_cache + ), patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new_callable=AsyncMock, + return_value=fake_key_data, + ) as mock_gen: + result = await create_jwt_client(data=data, user_api_key_dict=_make_admin_auth()) + + mock_gen.assert_called_once() + call_kwargs = mock_gen.call_args.kwargs + assert call_kwargs["allowed_routes"] == ["llm_api_routes"] + assert call_kwargs["metadata"]["jwt_bound"] is True + assert call_kwargs["metadata"]["jwt_claim_name"] == "sub" + + mock_prisma.db.litellm_jwtkeymapping.create.assert_called_once() + assert isinstance(result, JWTKeyMappingResponse) + + +@pytest.mark.asyncio +async def test_jwt_client_new_cache_invalidated(): + """/jwt_client/new must invalidate the cache for the new mapping.""" + from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import ( + create_jwt_client, + ) + from litellm.proxy._types import CreateJWTClientRequest + + mock_prisma = _mock_prisma() + mock_prisma.db.litellm_jwtkeymapping.create.return_value = _mock_mapping( + claim_name="sub", claim_value="svc-b" + ) + mock_cache = MagicMock() + mock_cache.async_delete_cache = AsyncMock() + + data = CreateJWTClientRequest(jwt_claim_name="sub", jwt_claim_value="svc-b") + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_cache + ), patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new_callable=AsyncMock, + return_value={"token": "sk-xyz"}, + ): + await create_jwt_client(data=data, user_api_key_dict=_make_admin_auth()) + + mock_cache.async_delete_cache.assert_called_once() + cache_key_arg = mock_cache.async_delete_cache.call_args.args[0] + assert "sub" in cache_key_arg and "svc-b" in cache_key_arg + + +@pytest.mark.asyncio +async def test_jwt_client_new_non_admin_rejected(): + """/jwt_client/new must reject non-admin users with 403.""" + from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import ( + create_jwt_client, + ) + from litellm.proxy._types import CreateJWTClientRequest + + data = CreateJWTClientRequest(jwt_claim_name="sub", jwt_claim_value="svc-c") + + with pytest.raises(HTTPException) as exc_info: + await create_jwt_client(data=data, user_api_key_dict=_make_non_admin_auth()) + assert exc_info.value.status_code == 403 + + +# ────────────────────────────────────────────── +# Block 4: Gap #4 — unregistered_jwt_client_behavior +# ────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_reject_behavior_raises_403_for_unknown_jwt(): + """'reject' mode should raise 403 when no mapping exists.""" + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + unregistered_jwt_client_behavior="reject", + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) + + with pytest.raises(HTTPException) as exc_info: + await _resolve_jwt_to_virtual_key( + jwt_claims={"sub": "unknown-svc"}, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=DualCache(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_fallback_behavior_returns_none_for_unknown_jwt(): + """Default 'fallback_team_mapping' mode should return None for unknown clients.""" + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + # default: fallback_team_mapping + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) + + with patch( + "litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock + ): + result = await _resolve_jwt_to_virtual_key( + jwt_claims={"sub": "unknown-svc"}, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=DualCache(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + assert result is None + + +@pytest.mark.asyncio +async def test_auto_register_creates_key_on_first_request(): + """'auto_register' mode should create a key+mapping on first unknown JWT.""" + from litellm.proxy._types import JWTClientAutoRegisterDefaults + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + unregistered_jwt_client_behavior="auto_register", + auto_register_defaults=JWTClientAutoRegisterDefaults( + models=["gpt-4o-mini"], max_budget=5.0 + ), + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() + + mock_key_obj = UserAPIKeyAuth(token="sk-auto", team_id=None) + + with patch( + "litellm.proxy.auth.user_api_key_auth.get_key_object", + new_callable=AsyncMock, + return_value=mock_key_obj, + ), patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new_callable=AsyncMock, + return_value={"token": "sk-auto"}, + ) as mock_gen: + result = await _resolve_jwt_to_virtual_key( + jwt_claims={"sub": "new-svc"}, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=DualCache(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + + assert result is not None + mock_gen.assert_called_once() + call_kwargs = mock_gen.call_args.kwargs + assert call_kwargs["allowed_routes"] == ["llm_api_routes"] + assert call_kwargs["metadata"]["jwt_bound"] is True + + +# ────────────────────────────────────────────── +# Block 5: Gap #5 — issuer column +# ────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_get_jwt_key_mapping_object_passes_issuer_to_db(): + """get_jwt_key_mapping_object must include issuer in the DB where clause.""" + from litellm.proxy.auth.auth_checks import get_jwt_key_mapping_object + + prisma_client = MagicMock() + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) + + await get_jwt_key_mapping_object( + jwt_claim_name="sub", + jwt_claim_value="svc", + prisma_client=prisma_client, + issuer="https://idp1.example.com", + ) + + call_kwargs = prisma_client.db.litellm_jwtkeymapping.find_first.call_args.kwargs + assert call_kwargs["where"]["issuer"] == "https://idp1.example.com" + + +@pytest.mark.asyncio +async def test_resolve_extracts_issuer_from_jwt_claims(): + """_resolve_jwt_to_virtual_key must pass iss claim as issuer to DB lookup.""" + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(virtual_key_claim_field="sub") + + prisma_client = MagicMock() + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) + + with patch( + "litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock + ): + await _resolve_jwt_to_virtual_key( + jwt_claims={"sub": "svc-a", "iss": "https://idp2.example.com"}, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=DualCache(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + + call_kwargs = prisma_client.db.litellm_jwtkeymapping.find_first.call_args.kwargs + assert call_kwargs["where"]["issuer"] == "https://idp2.example.com" + + +@pytest.mark.asyncio +async def test_no_issuer_in_jwt_defaults_to_empty_string(): + """When JWT has no 'iss' field, issuer should default to empty string.""" + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(virtual_key_claim_field="sub") + + prisma_client = MagicMock() + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) + + with patch( + "litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock + ): + await _resolve_jwt_to_virtual_key( + jwt_claims={"sub": "svc-b"}, # no "iss" + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=DualCache(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + + call_kwargs = prisma_client.db.litellm_jwtkeymapping.find_first.call_args.kwargs + assert call_kwargs["where"]["issuer"] == "" + + +# ────────────────────────────────────────────── +# Block 6: Gap #6 — endpoints expose virtual key properties +# ────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_info_endpoint_returns_virtual_key_fields(): + """/jwt/key/mapping/info should return virtual key fields when the key row exists.""" + now = datetime.now(timezone.utc) + mock_prisma = _mock_prisma() + mock_prisma.db.litellm_jwtkeymapping.find_unique.return_value = _mock_mapping() + + key_row = MagicMock() + key_row.models = ["gpt-4o"] + key_row.max_budget = 50.0 + key_row.budget_duration = "30d" + key_row.tpm_limit = 10000 + key_row.rpm_limit = 100 + key_row.team_id = "team-1" + key_row.key_alias = "my-client" + key_row.spend = 2.5 + key_row.expires = now + mock_prisma.db.litellm_verificationtoken.find_first.return_value = key_row + mock_cache = AsyncMock() + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_cache + ): + result = await info_jwt_key_mapping( + id="mapping-1", user_api_key_dict=_make_admin_auth() + ) + + assert result.models == ["gpt-4o"] + assert result.max_budget == 50.0 + assert result.team_id == "team-1" + assert result.key_alias == "my-client" + assert result.spend == 2.5 + + +@pytest.mark.asyncio +async def test_info_endpoint_still_returns_mapping_fields(): + """/jwt/key/mapping/info response must still include mapping metadata.""" + mock_prisma = _mock_prisma() + mock_prisma.db.litellm_jwtkeymapping.find_unique.return_value = _mock_mapping( + claim_name="email", claim_value="user@corp.com" + ) + mock_prisma.db.litellm_verificationtoken.find_first.return_value = None + mock_cache = AsyncMock() + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.proxy_server.user_api_key_cache", mock_cache + ): + result = await info_jwt_key_mapping( + id="mapping-1", user_api_key_dict=_make_admin_auth() + ) + + assert result.jwt_claim_name == "email" + assert result.jwt_claim_value == "user@corp.com" + assert result.is_active is True