mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
feat(proxy): JWT-to-virtual-key mapping improvements (P0–P2)
- 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
This commit is contained in:
parent
5ebc7b53b4
commit
de23b3ec1d
8 changed files with 1050 additions and 90 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue