Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/auth/handle_jwt.py: 16%
1062 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Supports using JWT's for authenticating into the proxy.
4Currently only supports admin.
6JWT token must have 'litellm_proxy_admin' in scope.
7"""
9from __future__ import annotations
11import asyncio
12import fnmatch
13import hashlib
14import os
15import re
16import time
17from collections.abc import Awaitable, Callable, Collection, Mapping, Sequence
18from dataclasses import dataclass
19from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast
21import httpx
22import jwt
23from cryptography import x509
24from cryptography.hazmat.backends import default_backend
25from cryptography.hazmat.primitives import serialization
26from fastapi import HTTPException, status
27from jwt.api_jwk import PyJWK
28from typing_extensions import ReadOnly, TypedDict
30from litellm._logging import verbose_proxy_logger
31from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
32from litellm.llms.custom_httpx.httpx_handler import HTTPHandler
33from litellm.proxy._types import (
34 DEFAULT_JWKS_STALE_TTL,
35 RBAC_ROLES,
36 JWKKeyValue,
37 JWTAuthBuilderResult,
38 JWTIssuerConfig,
39 JWTKeyItem,
40 LiteLLM_EndUserTable,
41 LiteLLM_JWTAuth,
42 LiteLLM_OrganizationTable,
43 LiteLLM_TeamMembership,
44 LiteLLM_TeamTable,
45 LiteLLM_UserTable,
46 LitellmUserRoles,
47 Member,
48 ProxyErrorTypes,
49 ProxyException,
50 ScopeMapping,
51 Span,
52 TeamMemberAddRequest,
53 UserAPIKeyAuth,
54)
55from litellm.proxy.auth.auth_checks import can_team_access_model
56from litellm.proxy.auth.model_access_denied import (
57 ModelAccessDeniedHTTPException,
58 model_access_denied_client_message,
59)
60from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canonical_user_id
61from litellm.proxy.auth.route_checks import RouteChecks
62from litellm.proxy.auth.team_grants import team_grants, team_model_aliases
63from litellm.proxy.common_utils.user_api_key_cache import (
64 UserApiKeyCache,
65 get_management_object_ttl,
66)
67from litellm.proxy.utils import PrismaClient, ProxyLogging
68from litellm.repositories.user_repository import UserRepository
69from litellm.types.agents import AgentResponse
70from litellm.types.proxy.auth.auth_checks import UserNotFoundError
72from .auth_checks import (
73 TeamNotFoundError,
74 _allowed_routes_check,
75 allowed_routes_check,
76 get_actual_routes,
77 get_end_user_object,
78 get_org_object,
79 get_org_object_by_alias,
80 get_role_based_models,
81 get_role_based_routes,
82 get_team_membership,
83 get_team_object,
84 get_team_object_by_alias,
85 get_user_object,
86)
89class NoMatchingJWTPublicKeyError(Exception):
90 """Raised when a JWKS endpoint returns no key matching the requested ``kid``."""
93class JWKSUnreachableError(Exception):
94 """Raised when an IdP's JWKS / OIDC discovery endpoint is unreachable and no cached copy is left to fall back on."""
97JWKS_FETCH_ATTEMPTS: Final = 3
98JWKS_FETCH_RETRY_BACKOFF_SECONDS: Final = 0.25
99JWKS_UNREACHABLE_BACKOFF_SECONDS: Final = 30
100STALE_CACHE_KEY_PREFIX: Final = "litellm_stale_"
101STALE_WRITTEN_AT_CACHE_KEY_PREFIX: Final = "litellm_stale_written_at_"
102UNREACHABLE_CACHE_KEY_PREFIX: Final = "litellm_jwks_unreachable_"
104_CachedValueT = TypeVar("_CachedValueT", bound=JWKKeyValue | str)
107class _JWTAuthSettings(Protocol):
108 """The JWT auth settings block this handler reads back through ``getattr``, when one is configured."""
110 @property
111 def issuers(self) -> Sequence[JWTIssuerConfig] | None: ... 111 ↛ exitline 111 didn't return from function 'issuers' because
113 @property
114 def public_key_ttl(self) -> float: ... 114 ↛ exitline 114 didn't return from function 'public_key_ttl' because
116 @property
117 def public_key_stale_ttl(self) -> float: ... 117 ↛ exitline 117 didn't return from function 'public_key_stale_ttl' because
120class _OIDCDiscoveryBody(TypedDict, total=False):
121 """Decoded OIDC discovery document, read for the JWKS endpoint it advertises."""
123 jwks_uri: ReadOnly[str]
126class _OIDCDiscoveryResponse(Protocol):
127 """The discovery endpoint's HTTP response, read for the decoded document it carries."""
129 def json(self) -> _OIDCDiscoveryBody: ... 129 ↛ exitline 129 didn't return from function 'json' because
132class _UserInfoResponse(Protocol):
133 """The OIDC UserInfo endpoint's HTTP response, read for the identity document it carries."""
135 def json(self) -> dict[str, object]: ... 135 ↛ exitline 135 didn't return from function 'json' because
138@dataclass(frozen=True, slots=True)
139class JWTIdentity:
140 user_id: str | None
141 user_object: LiteLLM_UserTable | None
142 agent_id: str | None
145@dataclass(frozen=True, slots=True)
146class _JWTProvisioning:
147 user_id_upsert: bool
148 team_id_upsert: bool
151@dataclass(frozen=True, slots=True)
152class HeaderTeam:
153 header_value: str
154 team_id: str
157class AgentLookup(Protocol):
158 """The registered-agent lookups a JWT agent claim is matched against."""
160 def get_agent_by_id(self, agent_id: str) -> AgentResponse | None:
161 """The agent registered under ``agent_id``, if any."""
163 def get_agent_by_name(self, agent_name: str) -> AgentResponse | None:
164 """The agent registered under ``agent_name``, if any."""
167class _NoRegisteredAgents:
168 """The lookup in force until the proxy binds its agent registry: no agent is registered, so no claim matches."""
170 def get_agent_by_id(self, agent_id: str) -> None:
171 return None
173 def get_agent_by_name(self, agent_name: str) -> None:
174 return None
177def _discovery_document(response: _OIDCDiscoveryResponse) -> _OIDCDiscoveryBody:
178 """Decode an OIDC discovery response body."""
179 return response.json()
182def _userinfo_document(response: _UserInfoResponse) -> dict[str, object]:
183 """Decode an OIDC UserInfo response body into its JSON object form."""
184 return response.json()
187def jwks_unavailable_exception(error: JWKSUnreachableError) -> ProxyException:
188 return ProxyException(
189 message=(
190 "Service Unavailable, the identity provider's JWKS endpoint is temporarily "
191 f"unreachable, so the JWT signature could not be verified. Please retry shortly. Error: {error}"
192 ),
193 type=ProxyErrorTypes.auth_provider_unavailable,
194 param="None",
195 code=status.HTTP_503_SERVICE_UNAVAILABLE,
196 )
199class JWTHandler:
200 """
201 - treat the sub id passed in as the user id
202 - return an error if id making request doesn't exist in proxy user table
203 - track spend against the user id
204 - if role="litellm_proxy_user" -> allow making calls + info. Can not edit budgets
205 """
207 prisma_client: PrismaClient | None
208 user_api_key_cache: UserApiKeyCache
209 # Supported algos: https://pyjwt.readthedocs.io/en/stable/algorithms.html
210 # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret
211 # the key in different ways (e.g. HS* and RS*)."
212 SUPPORTED_JWT_ALGORITHMS = [
213 "RS256",
214 "RS384",
215 "RS512",
216 "PS256",
217 "PS384",
218 "PS512",
219 "ES256",
220 "ES384",
221 "ES512",
222 "EdDSA",
223 ]
224 LITELLM_JWT_ISSUER_CLAIM = "_litellm_jwt_issuer"
225 LITELLM_USER_ID_CLAIM = "_litellm_user_id"
226 LITELLM_USER_EMAIL_CLAIM = "_litellm_user_email"
227 LITELLM_TEAM_ID_CLAIM = "_litellm_team_id"
228 LITELLM_TEAM_IDS_CLAIM = "_litellm_team_ids"
229 LITELLM_ORG_ID_CLAIM = "_litellm_org_id"
230 LITELLM_END_USER_ID_CLAIM = "_litellm_end_user_id"
231 LITELLM_INTERNAL_CLAIMS = (
232 LITELLM_JWT_ISSUER_CLAIM,
233 LITELLM_USER_ID_CLAIM,
234 LITELLM_USER_EMAIL_CLAIM,
235 LITELLM_TEAM_ID_CLAIM,
236 LITELLM_TEAM_IDS_CLAIM,
237 LITELLM_ORG_ID_CLAIM,
238 LITELLM_END_USER_ID_CLAIM,
239 )
241 def __init__(
242 self,
243 ) -> None:
244 self.http_handler = HTTPHandler()
245 self.leeway = 0
246 # Per-cache-key locks so a TTL lapse triggers one refresh instead of one per in-flight request.
247 self._refresh_locks: dict[str, asyncio.Lock] = {} # mutable-ok: lock registry, keyed by JWKS url
248 self.agent_lookup: AgentLookup = _NoRegisteredAgents()
250 def bind_agent_lookup(self, agent_lookup: AgentLookup) -> None:
251 self.agent_lookup = agent_lookup
253 def update_environment(
254 self,
255 prisma_client: PrismaClient | None,
256 user_api_key_cache: UserApiKeyCache,
257 litellm_jwtauth: LiteLLM_JWTAuth,
258 leeway: int = 0,
259 ) -> None:
260 self.prisma_client = prisma_client
261 self.user_api_key_cache = user_api_key_cache
262 self.litellm_jwtauth = litellm_jwtauth
263 self.leeway = leeway
265 @staticmethod
266 def is_jwt(token: str | None) -> bool:
267 if token is None: 267 ↛ 268line 267 didn't jump to line 268 because the condition on line 267 was never true
268 return False
269 parts: Final = token.split(".")
270 return len(parts) == 3
272 @staticmethod
273 def get_unverified_claims(token: str) -> dict | None:
274 """
275 Decode JWT claims without signature verification.
276 Used for routing decisions before selecting validation path.
277 """
278 if not JWTHandler.is_jwt(token):
279 return None
281 try:
282 claims: Final = jwt.decode(
283 token,
284 options={"verify_signature": False, "verify_aud": False},
285 algorithms=JWTHandler.SUPPORTED_JWT_ALGORITHMS,
286 )
287 if isinstance(claims, dict):
288 return claims
289 return None
290 except Exception as e:
291 verbose_proxy_logger.debug("Failed to decode unverified JWT claims for routing: %s", e)
292 return None
294 def _rbac_role_from_role_mapping(self, token: dict) -> RBAC_ROLES | None:
295 """
296 Returns the RBAC role the token 'belongs' to based on role mappings.
298 Args:
299 token (dict): The JWT token containing role information
301 Returns:
302 Optional[RBAC_ROLES]: The mapped internal RBAC role if a mapping exists,
303 None otherwise
305 Note:
306 The function handles both single string roles and lists of roles from the JWT.
307 If multiple mappings match the JWT roles, the first matching mapping is returned.
308 """
309 if self.litellm_jwtauth.role_mappings is None:
310 return None
312 jwt_role: Final = self.get_jwt_role(token=token, default_value=None)
313 if not jwt_role:
314 return None
316 jwt_role_set: Final = set(jwt_role)
318 for role_mapping in self.litellm_jwtauth.role_mappings:
319 # Check if the mapping role matches any of the JWT roles
320 if role_mapping.role in jwt_role_set:
321 return role_mapping.internal_role
323 return None
325 def get_rbac_role(self, token: dict) -> RBAC_ROLES | None:
326 """
327 Returns the RBAC role the token 'belongs' to.
329 RBAC roles allowed to make requests:
330 - PROXY_ADMIN: can make requests to all routes
331 - TEAM: can make requests to routes associated with a team
332 - INTERNAL_USER: can make requests to routes associated with a user
334 Resolves: https://github.com/BerriAI/litellm/issues/6793
336 Returns:
337 - PROXY_ADMIN: if token is admin
338 - TEAM: if token is associated with a team
339 - INTERNAL_USER: if token is associated with a user
340 - None: if token is not associated with a team or user
341 """
342 scopes: Final = self.get_scopes(token=token)
343 is_admin: Final = self.is_admin(scopes=scopes)
344 user_roles: Final = self.get_user_roles(token=token, default_value=None)
346 if is_admin:
347 return LitellmUserRoles.PROXY_ADMIN
348 elif self.get_team_id(token=token, default_value=None) is not None:
349 return LitellmUserRoles.TEAM
350 elif (
351 self.get_user_id(token=token, default_value=None) is not None
352 or user_roles is not None
353 and self.is_allowed_user_role(user_roles=user_roles)
354 ):
355 return LitellmUserRoles.INTERNAL_USER
356 elif rbac_role := self._rbac_role_from_role_mapping(token=token):
357 return rbac_role
359 return None
361 def is_admin(self, scopes: list) -> bool:
362 if self.litellm_jwtauth.admin_jwt_scope in scopes:
363 return True
364 return False
366 def _is_trusted_issuer_normalized_token(self, token: dict) -> bool:
367 issuer: Final = token.get(self.LITELLM_JWT_ISSUER_CLAIM)
368 if not isinstance(issuer, str) or not issuer:
369 return False
371 litellm_jwtauth: Final = getattr(self, "litellm_jwtauth", None)
372 issuer_configs: Final = getattr(litellm_jwtauth, "issuers", None) or []
373 return any(issuer_config.issuer == issuer for issuer_config in issuer_configs)
375 def _has_trusted_issuer_normalized_claim(self, token: dict, claim: str) -> bool:
376 return self._is_trusted_issuer_normalized_token(token=token) and claim in token
378 def get_team_ids_from_jwt(self, token: dict) -> list[str]:
379 if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_TEAM_IDS_CLAIM):
380 issuer_team_ids: Final = token.get(self.LITELLM_TEAM_IDS_CLAIM)
381 if isinstance(issuer_team_ids, list):
382 return issuer_team_ids
383 if isinstance(issuer_team_ids, str):
384 return [issuer_team_ids]
385 # Issuer-scoped claim exists but has an unexpected type
386 # (e.g. int/dict from an unusual upstream mapping). Don't silently
387 # fall through to the global ``team_ids_jwt_field`` path — that
388 # would read a semantically unrelated claim on the same token.
389 return []
391 if self.litellm_jwtauth.team_ids_jwt_field is not None:
392 team_ids: Final[list[str] | None] = get_nested_value(
393 data=token,
394 key_path=self.litellm_jwtauth.team_ids_jwt_field,
395 default=[],
396 )
397 return team_ids or []
399 return []
401 def get_all_jwt_team_ids(self, token: dict) -> list[str]:
402 """
403 Return team IDs from both the plural ``team_ids_jwt_field`` and the
404 singular ``team_id_jwt_field`` claim (string or list of strings), as a
405 deduplicated list preserving plural-first order.
407 Membership-reconciliation paths (SSO callback, JWT-bearer sync) need
408 to consider both claim shapes. Reading only the plural field — as
409 callers historically did — silently dropped users whose IdP populates
410 the singular field, which is what Okta and Auth0 default to when a
411 user has a single primary team.
413 This intentionally does NOT consult ``team_id_default``: that fallback
414 is a property of how the JWT-bearer auth flow resolves a single
415 request-bound team, not of the token's claims. Callers that want the
416 default-team behavior should still go through ``get_team_id``.
417 """
418 team_ids: Final[list[str]] = list(self.get_team_ids_from_jwt(token))
419 singular: Any = None
420 if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_TEAM_ID_CLAIM):
421 singular = token.get(self.LITELLM_TEAM_ID_CLAIM)
422 elif self.litellm_jwtauth.team_id_jwt_field is not None:
423 singular = get_nested_value(
424 data=token,
425 key_path=self.litellm_jwtauth.team_id_jwt_field,
426 default=None,
427 )
428 if singular is not None:
429 if isinstance(singular, list):
430 for item in singular:
431 if item is None:
432 continue
433 sid = str(item)
434 if sid and sid not in team_ids:
435 team_ids.append(sid)
436 elif singular and str(singular) not in team_ids:
437 team_ids.append(str(singular))
438 return team_ids
440 def get_end_user_id(self, token: dict, default_value: str | None) -> str | None:
441 if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_END_USER_ID_CLAIM):
442 return token.get(self.LITELLM_END_USER_ID_CLAIM)
444 try:
445 if self.litellm_jwtauth.end_user_id_jwt_field is not None:
446 user_id = get_nested_value(
447 data=token,
448 key_path=self.litellm_jwtauth.end_user_id_jwt_field,
449 default=default_value,
450 )
451 else:
452 user_id = None
453 except KeyError:
454 user_id = default_value
456 return user_id
458 def is_required_team_id(self) -> bool:
459 """
460 Returns:
461 - True: if 'team_id_jwt_field' or 'team_alias_jwt_field' is set
462 - False: if neither is set
463 """
464 if self.litellm_jwtauth.team_id_jwt_field is None and self.litellm_jwtauth.team_alias_jwt_field is None:
465 return False
466 return True
468 def is_enforced_email_domain(self) -> bool:
469 """
470 Returns:
471 - True: if 'user_allowed_email_domain' is set
472 - False: if 'user_allowed_email_domain' is None
473 """
475 if self.litellm_jwtauth.user_allowed_email_domain is not None and isinstance(
476 self.litellm_jwtauth.user_allowed_email_domain, str
477 ):
478 return True
479 return False
481 def get_team_id(self, token: dict, default_value: str | None) -> str | None:
482 if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_TEAM_ID_CLAIM):
483 team_id = token.get(self.LITELLM_TEAM_ID_CLAIM)
484 if isinstance(team_id, list):
485 return team_id[0] if team_id else default_value
486 return team_id
488 try:
489 if self.litellm_jwtauth.team_id_jwt_field is not None:
490 # Use a sentinel value to detect if the path actually exists
491 sentinel: Final = object()
492 team_id = get_nested_value(
493 data=token,
494 key_path=self.litellm_jwtauth.team_id_jwt_field,
495 default=sentinel,
496 )
497 if team_id is sentinel:
498 # Path doesn't exist, use team_id_default if available
499 if self.litellm_jwtauth.team_id_default is not None:
500 return self.litellm_jwtauth.team_id_default
501 else:
502 return default_value
503 # AAD and other IdPs often send roles/groups as a list of strings.
504 # team_id_jwt_field is singular, so take the first element when a list
505 # is returned. This avoids "unhashable type: 'list'" errors downstream.
506 if isinstance(team_id, list):
507 if not team_id:
508 return default_value
509 verbose_proxy_logger.debug(
510 "JWT Auth: team_id_jwt_field '%s' returned a list %s; using first element '%s' automatically.",
511 self.litellm_jwtauth.team_id_jwt_field,
512 team_id,
513 team_id[0],
514 )
515 team_id = team_id[0]
516 return team_id
517 elif self.litellm_jwtauth.team_id_default is not None:
518 team_id = self.litellm_jwtauth.team_id_default
519 else:
520 team_id = None
521 except KeyError:
522 team_id = default_value
523 return team_id
525 def get_team_alias(self, token: dict, default_value: str | None) -> str | None:
526 """
527 Extract team name/alias from JWT token using the configured team_alias_jwt_field.
529 Args:
530 token: The decoded JWT token dictionary
531 default_value: Default value to return if field not found
533 Returns:
534 The team alias from the token, or default_value if not found
535 """
536 try:
537 if self.litellm_jwtauth.team_alias_jwt_field is not None:
538 team_alias = get_nested_value(
539 data=token,
540 key_path=self.litellm_jwtauth.team_alias_jwt_field,
541 default=default_value,
542 )
543 return team_alias
544 else:
545 team_alias = None
546 except KeyError:
547 team_alias = default_value
548 return team_alias
550 def is_upsert_user_id(self, valid_user_email: bool | None = None) -> bool:
551 """
552 Returns:
553 - True: if 'user_id_upsert' is set AND valid_user_email is not False
554 - False: if not
555 """
556 if valid_user_email is False:
557 return False
558 return self.litellm_jwtauth.user_id_upsert
560 def get_user_id(self, token: dict, default_value: str | None) -> str | None:
561 if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_USER_ID_CLAIM):
562 return token.get(self.LITELLM_USER_ID_CLAIM)
564 try:
565 if self.litellm_jwtauth.user_id_jwt_field is not None:
566 user_id = get_nested_value(
567 data=token,
568 key_path=self.litellm_jwtauth.user_id_jwt_field,
569 default=default_value,
570 )
571 else:
572 user_id = default_value
573 except KeyError:
574 user_id = default_value
575 return user_id
577 def get_user_roles(self, token: dict, default_value: list[str] | None) -> list[str] | None:
578 """
579 Returns the user role from the token.
581 Set via 'user_roles_jwt_field' in the config.
582 """
583 try:
584 if self.litellm_jwtauth.user_roles_jwt_field is not None:
585 user_roles = get_nested_value(
586 data=token,
587 key_path=self.litellm_jwtauth.user_roles_jwt_field,
588 default=default_value,
589 )
590 else:
591 user_roles = default_value
592 except KeyError:
593 user_roles = default_value
594 return user_roles
596 def map_jwt_role_to_litellm_role(self, token: dict) -> LitellmUserRoles | None:
597 """Map roles from JWT to LiteLLM user roles"""
598 if not self.litellm_jwtauth.jwt_litellm_role_map:
599 return None
601 jwt_roles: Final = self.get_jwt_role(token=token, default_value=[])
602 if not jwt_roles:
603 return None
605 for mapping in self.litellm_jwtauth.jwt_litellm_role_map:
606 for role in jwt_roles:
607 if fnmatch.fnmatch(role, mapping.jwt_role):
608 return mapping.litellm_role
609 return None
611 def get_jwt_role(self, token: dict, default_value: list[str] | None) -> list[str] | None:
612 """
613 Generic implementation of `get_user_roles` that can be used for both user and team roles.
615 Returns the jwt role from the token.
617 Set via 'roles_jwt_field' in the config.
618 """
619 try:
620 if self.litellm_jwtauth.roles_jwt_field is not None:
621 user_roles = get_nested_value(
622 data=token,
623 key_path=self.litellm_jwtauth.roles_jwt_field,
624 default=default_value,
625 )
626 else:
627 user_roles = default_value
628 except KeyError:
629 user_roles = default_value
630 return user_roles
632 def is_allowed_user_role(self, user_roles: list[str] | None) -> bool:
633 """
634 Returns the user role from the token.
636 Set via 'user_allowed_roles' in the config.
637 """
638 if (
639 user_roles is not None
640 and self.litellm_jwtauth.user_allowed_roles is not None
641 and any(role in self.litellm_jwtauth.user_allowed_roles for role in user_roles)
642 ):
643 return True
644 return False
646 def get_user_email(self, token: dict, default_value: str | None) -> str | None:
647 if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_USER_EMAIL_CLAIM):
648 return token.get(self.LITELLM_USER_EMAIL_CLAIM)
650 try:
651 if self.litellm_jwtauth.user_email_jwt_field is not None:
652 user_email = get_nested_value(
653 data=token,
654 key_path=self.litellm_jwtauth.user_email_jwt_field,
655 default=default_value,
656 )
657 else:
658 user_email = None
659 except KeyError:
660 user_email = default_value
661 return user_email
663 def get_object_id(self, token: dict, default_value: str | None) -> str | None:
664 try:
665 if self.litellm_jwtauth.object_id_jwt_field is not None:
666 object_id = get_nested_value(
667 data=token,
668 key_path=self.litellm_jwtauth.object_id_jwt_field,
669 default=default_value,
670 )
671 else:
672 object_id = default_value
673 except KeyError:
674 object_id = default_value
675 return object_id
677 def get_agent_claim(self, token: Mapping[str, object]) -> str | None:
678 if self.litellm_jwtauth.agent_id_jwt_field is None:
679 return None
680 claim: Final[object] = get_nested_value(data=token, key_path=self.litellm_jwtauth.agent_id_jwt_field)
681 return claim if isinstance(claim, str) and claim else None
683 def get_org_id(self, token: dict, default_value: str | None) -> str | None:
684 if self._has_trusted_issuer_normalized_claim(token=token, claim=self.LITELLM_ORG_ID_CLAIM):
685 return token.get(self.LITELLM_ORG_ID_CLAIM)
687 try:
688 if self.litellm_jwtauth.org_id_jwt_field is not None:
689 org_id = get_nested_value(
690 data=token,
691 key_path=self.litellm_jwtauth.org_id_jwt_field,
692 default=default_value,
693 )
694 else:
695 org_id = None
696 except KeyError:
697 org_id = default_value
698 return org_id
700 def get_org_alias(self, token: dict, default_value: str | None) -> str | None:
701 """
702 Extract organization name/alias from JWT token using the configured org_alias_jwt_field.
704 Args:
705 token: The decoded JWT token dictionary
706 default_value: Default value to return if field not found
708 Returns:
709 The organization alias from the token, or default_value if not found
710 """
711 try:
712 if self.litellm_jwtauth.org_alias_jwt_field is not None:
713 org_alias = get_nested_value(
714 data=token,
715 key_path=self.litellm_jwtauth.org_alias_jwt_field,
716 default=default_value,
717 )
718 return org_alias
719 else:
720 org_alias = None
721 except KeyError:
722 org_alias = default_value
723 return org_alias
725 def get_scopes(self, token: dict) -> list[str]:
726 try:
727 if isinstance(token["scope"], str):
728 # Assuming the scopes are stored in 'scope' claim and are space-separated
729 scopes = token["scope"].split()
730 elif isinstance(token["scope"], list):
731 scopes = token["scope"]
732 else:
733 raise Exception(f"Unmapped scope type - {type(token['scope'])}. Supported types - list, str.")
734 except KeyError:
735 scopes = []
736 return scopes
738 async def _resolve_jwks_url(self, url: str) -> str:
739 """
740 If url points to an OIDC discovery document (*.well-known/openid-configuration),
741 fetch it and return the jwks_uri contained within. Otherwise return url unchanged.
742 This lets JWT_PUBLIC_KEY_URL be set to a well-known discovery endpoint instead of
743 requiring operators to manually find the JWKS URL.
744 """
745 if ".well-known/openid-configuration" not in url:
746 return url
748 return await self._cached_with_stale_fallback(
749 cache_key=f"litellm_oidc_discovery_{url}",
750 ttl=self._get_public_key_cache_ttl(),
751 refresh=lambda: self._fetch_jwks_uri_from_discovery(url),
752 log_context="an OIDC discovery lookup",
753 )
755 async def _get_with_transient_retries(self, url: str) -> httpx.Response:
756 """GET ``url``, retrying transport failures so one IdP blip does not fail the request."""
757 for attempt in range(1, JWKS_FETCH_ATTEMPTS):
758 try:
759 return await self.http_handler.get(url)
760 except httpx.TransportError as e:
761 verbose_proxy_logger.warning(
762 "JWT Auth: %s fetching %s (attempt %s/%s), retrying: %s",
763 type(e).__name__,
764 url,
765 attempt,
766 JWKS_FETCH_ATTEMPTS,
767 e,
768 )
769 await asyncio.sleep(JWKS_FETCH_RETRY_BACKOFF_SECONDS * attempt)
771 try:
772 return await self.http_handler.get(url)
773 except httpx.TransportError as e:
774 raise JWKSUnreachableError(f"{type(e).__name__} fetching {url} after {JWKS_FETCH_ATTEMPTS} attempts") from e
776 async def _get_cached_value(self, cache_key: str) -> _CachedValueT | None:
777 cached: Final = await self.user_api_key_cache.async_get_cache(cache_key)
778 return cast("_CachedValueT | None", cached) # cast-ok: cache reads are untyped
780 async def _get_cached_timestamp(self, cache_key: str) -> float | None:
781 cached: Final = await self.user_api_key_cache.async_get_cache(cache_key)
782 # A JSON round-trip through Redis hands a whole-number epoch back as an int.
783 return float(cached) if isinstance(cached, (int, float)) else None
785 async def _put_cached_value(self, cache_key: str, value: JWKKeyValue | str | float, ttl: float) -> None:
786 await self.user_api_key_cache.async_set_cache(key=cache_key, value=value, ttl=ttl)
788 async def _cached_with_stale_fallback(
789 self,
790 cache_key: str,
791 ttl: float,
792 refresh: Callable[[], Awaitable[_CachedValueT]],
793 log_context: str,
794 ) -> _CachedValueT:
795 """Read ``cache_key``, refreshing it through a single-flight lock on a miss."""
796 cached: Final[_CachedValueT | None] = await self._get_cached_value(cache_key)
797 if cached is not None:
798 return cached
800 lock: Final = self._refresh_locks.setdefault(cache_key, asyncio.Lock())
801 async with lock:
802 cached_after_lock: Final[_CachedValueT | None] = await self._get_cached_value(cache_key)
803 if cached_after_lock is not None:
804 return cached_after_lock
805 return await self._refresh_or_serve_stale(
806 cache_key=cache_key, ttl=ttl, refresh=refresh, log_context=log_context
807 )
809 async def _refresh_or_serve_stale(
810 self,
811 cache_key: str,
812 ttl: float,
813 refresh: Callable[[], Awaitable[_CachedValueT]],
814 log_context: str,
815 ) -> _CachedValueT:
816 """Refresh ``cache_key`` from the IdP, falling back to the last-known-good copy when it is unreachable.
818 Signing keys rotate rarely, so a last-known-good key beats failing authentication during an IdP blip.
819 How long a key the IdP has since removed stays trusted is bounded by ``public_key_ttl`` +
820 ``public_key_stale_ttl`` measured from when the copy was taken, and that bound is enforced here on every
821 read rather than baked into the cache entry's own expiry. An operator who lowers ``public_key_stale_ttl``,
822 or sets it to 0 to fail closed, is usually doing it mid-incident, and a copy written under the old longer
823 setting would otherwise stay servable until it aged out on its own. A copy whose write time cannot be
824 established is not servable, so the bound cannot be dodged by losing the timestamp.
825 """
826 stale_ttl: Final = self._get_public_key_stale_ttl()
827 outcome: Final = await self._refresh_or_record_outage(
828 cache_key=cache_key, ttl=ttl, stale_ttl=stale_ttl, refresh=refresh
829 )
830 if not isinstance(outcome, JWKSUnreachableError):
831 return outcome
832 if stale_ttl <= 0:
833 raise outcome
835 stale: Final[_CachedValueT | None] = await self._get_cached_value(f"{STALE_CACHE_KEY_PREFIX}{cache_key}")
836 age: Final = await self._stale_copy_age(cache_key)
837 lifetime: Final = ttl + stale_ttl
838 if stale is None or age is None or age > lifetime:
839 raise outcome
840 verbose_proxy_logger.warning(
841 "JWT Auth: identity provider unreachable, authenticating %s against a stale JWKS copy of %s "
842 "(last refreshed %.0fs ago, stops being trusted in %.0fs). Refresh failed: %s",
843 log_context,
844 cache_key,
845 age,
846 max(lifetime - age, 0),
847 outcome,
848 )
849 return stale
851 async def _stale_copy_age(self, cache_key: str) -> float | None:
852 written_at: Final = await self._get_cached_timestamp(f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{cache_key}")
853 return None if written_at is None else time.time() - written_at
855 async def _refresh_or_record_outage(
856 self,
857 cache_key: str,
858 ttl: float,
859 stale_ttl: float,
860 refresh: Callable[[], Awaitable[_CachedValueT]],
861 ) -> _CachedValueT | JWKSUnreachableError:
862 """Refresh ``cache_key``, returning the outage as a value rather than raising it.
864 A failed refresh is remembered for ``JWKS_UNREACHABLE_BACKOFF_SECONDS`` so a sustained outage costs one
865 fetch per window instead of one per request serialised behind the refresh lock.
866 """
867 unreachable_cache_key: Final = f"{UNREACHABLE_CACHE_KEY_PREFIX}{cache_key}"
868 recent_failure: Final[str | None] = await self._get_cached_value(unreachable_cache_key)
869 if recent_failure is not None:
870 return JWKSUnreachableError(recent_failure)
872 try:
873 refreshed: Final = await refresh()
874 except JWKSUnreachableError as e:
875 await self._put_cached_value(
876 cache_key=unreachable_cache_key, value=str(e), ttl=JWKS_UNREACHABLE_BACKOFF_SECONDS
877 )
878 return e
880 await self._put_cached_value(cache_key=cache_key, value=refreshed, ttl=ttl)
881 if stale_ttl > 0:
882 await self._put_cached_value(
883 cache_key=f"{STALE_CACHE_KEY_PREFIX}{cache_key}", value=refreshed, ttl=ttl + stale_ttl
884 )
885 await self._put_cached_value(
886 cache_key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{cache_key}", value=time.time(), ttl=ttl + stale_ttl
887 )
888 return refreshed
890 async def _fetch_jwks_uri_from_discovery(self, url: str) -> str:
891 verbose_proxy_logger.debug("JWT Auth: Fetching OIDC discovery document from %s", url)
892 response: Final = await self._get_with_transient_retries(url)
893 if response.status_code != 200:
894 raise Exception(
895 f"JWT Auth: OIDC discovery endpoint {url} returned status {response.status_code}: {response.text}"
896 )
897 try:
898 discovery: Final = _discovery_document(response)
899 except Exception as e:
900 raise Exception(f"JWT Auth: Failed to parse OIDC discovery document at {url}: {e}")
902 jwks_uri: Final = discovery.get("jwks_uri")
903 if not jwks_uri:
904 raise Exception(f"JWT Auth: OIDC discovery document at {url} does not contain a 'jwks_uri' field.")
906 verbose_proxy_logger.debug("JWT Auth: Resolved OIDC discovery %s -> jwks_uri=%s", url, jwks_uri)
907 return jwks_uri
909 def _get_public_key_cache_ttl(self) -> float:
910 litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
911 if litellm_jwtauth is None:
912 return 600
913 return litellm_jwtauth.public_key_ttl
915 def _get_public_key_stale_ttl(self) -> float:
916 litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
917 if litellm_jwtauth is None:
918 return DEFAULT_JWKS_STALE_TTL
919 return litellm_jwtauth.public_key_stale_ttl
921 async def _fetch_jwks_keys(self, resolved_jwks_url: str) -> JWKKeyValue:
922 response: Final = await self._get_with_transient_retries(resolved_jwks_url)
923 if response.status_code != 200:
924 raise Exception(
925 f"JWT Auth: JWKS endpoint {resolved_jwks_url} returned status {response.status_code}: {response.text}"
926 )
928 try:
929 response_json: Final = response.json()
930 except Exception as e:
931 verbose_proxy_logger.error("Error parsing response: %s. Original Response: %s", e, response.text)
932 raise Exception(f"Error parsing response: {e}. Check server logs for original response.")
934 keys: Final = response_json["keys"] if "keys" in response_json else response_json
935 return cast(JWKKeyValue, keys) # cast-ok: JWTKeyItem declares only `kid`, validating would drop key material
937 async def _get_public_key_from_jwks_url(self, jwks_url: str, kid: str | None) -> dict:
938 resolved_jwks_url: Final = await self._resolve_jwks_url(jwks_url)
939 keys: Final = await self._cached_with_stale_fallback(
940 cache_key=f"litellm_jwt_auth_keys_{resolved_jwks_url}",
941 ttl=self._get_public_key_cache_ttl(),
942 refresh=lambda: self._fetch_jwks_keys(resolved_jwks_url),
943 log_context=f"kid={kid}",
944 )
946 public_key: Final = self.parse_keys(keys=keys, kid=kid)
947 if public_key is not None:
948 return cast(dict, public_key)
950 raise NoMatchingJWTPublicKeyError(f"No matching public key found. keys={resolved_jwks_url}, kid={kid}")
952 async def get_public_key(self, kid: str | None) -> dict:
953 keys_url: Final = os.getenv("JWT_PUBLIC_KEY_URL")
955 if keys_url is None:
956 raise Exception("Missing JWT Public Key URL from environment.")
958 keys_url_list: Final = [url.strip() for url in keys_url.split(",") if url.strip()]
960 for key_url in keys_url_list:
961 try:
962 return await self._get_public_key_from_jwks_url(jwks_url=key_url, kid=kid)
963 except NoMatchingJWTPublicKeyError as e:
964 verbose_proxy_logger.debug("JWT Auth: No matching public key found at %s: %s", key_url, e)
965 except JWKSUnreachableError as e:
966 verbose_proxy_logger.error("JWT Auth: JWKS endpoint %s unreachable: %s", key_url, e)
967 raise jwks_unavailable_exception(e) from e
969 raise NoMatchingJWTPublicKeyError(f"No matching public key found. keys={keys_url_list}, kid={kid}")
971 def parse_keys(self, keys: JWKKeyValue, kid: str | None) -> JWTKeyItem | None:
972 public_key: JWTKeyItem | None = None
973 if len(keys) == 1:
974 if isinstance(keys, dict) and (keys.get("kid", None) == kid or kid is None):
975 public_key = keys
976 elif isinstance(keys, list) and (keys[0].get("kid", None) == kid or kid is None):
977 public_key = keys[0]
978 elif len(keys) > 1:
979 for key in keys:
980 if isinstance(key, dict):
981 key_kid = key.get("kid", None)
982 else:
983 key_kid = None
984 if kid is not None and isinstance(key, dict) and key_kid is not None and key_kid == kid:
985 public_key = key
987 return public_key
989 def is_allowed_domain(self, user_email: str) -> bool:
990 if self.litellm_jwtauth.user_allowed_email_domain is None:
991 return True
993 email_domain: Final = user_email.split("@")[-1] # Extract domain from email
994 if email_domain == self.litellm_jwtauth.user_allowed_email_domain:
995 return True
996 else:
997 return False
999 async def get_oidc_userinfo(self, token: str) -> dict:
1000 """
1001 Fetch user information from OIDC UserInfo endpoint.
1003 This follows the OpenID Connect protocol where an access token
1004 is sent to the identity provider's UserInfo endpoint to retrieve
1005 user identity information.
1007 Args:
1008 token: The access token to use for authentication
1010 Returns:
1011 dict: User information from the UserInfo endpoint
1013 Raises:
1014 Exception: If UserInfo endpoint is not configured or request fails
1015 """
1016 if not self.litellm_jwtauth.oidc_userinfo_endpoint:
1017 raise Exception("OIDC UserInfo endpoint not configured. Set 'oidc_userinfo_endpoint' in JWT auth config.")
1019 # Check cache first
1020 cache_key: Final = f"oidc_userinfo_{hashlib.sha256(token.encode()).hexdigest()}"
1021 cached_userinfo: Final = await self.user_api_key_cache.async_get_cache(cache_key)
1023 if cached_userinfo is not None:
1024 verbose_proxy_logger.debug("Returning cached OIDC UserInfo")
1025 return cached_userinfo
1027 verbose_proxy_logger.debug("Calling OIDC UserInfo endpoint: %s", self.litellm_jwtauth.oidc_userinfo_endpoint)
1029 try:
1030 # Call the UserInfo endpoint with the access token
1031 response: Final = await self.http_handler.get(
1032 url=self.litellm_jwtauth.oidc_userinfo_endpoint,
1033 headers={
1034 "Authorization": f"Bearer {token}",
1035 "Accept": "application/json",
1036 },
1037 )
1039 if response.status_code != 200:
1040 raise Exception(f"OIDC UserInfo endpoint returned status {response.status_code}: {response.text}")
1042 userinfo: Final = _userinfo_document(response)
1043 verbose_proxy_logger.debug("Received OIDC UserInfo: %s", userinfo)
1045 # Cache the userinfo response
1046 await self.user_api_key_cache.async_set_cache(
1047 key=cache_key,
1048 value=userinfo,
1049 ttl=self.litellm_jwtauth.oidc_userinfo_cache_ttl,
1050 )
1052 return userinfo
1054 except Exception as e:
1055 verbose_proxy_logger.error("Error fetching OIDC UserInfo: %s", e)
1056 raise Exception(f"Failed to fetch OIDC UserInfo: {e}")
1058 _unscoped_jwt_warning_emitted = False
1060 @classmethod
1061 def _build_decode_kwargs(cls) -> dict:
1062 """Build the audience/issuer/options kwargs for ``jwt.decode``.
1064 Setting ``JWT_AUDIENCE`` (and optionally ``JWT_ISSUER``) turns on the
1065 corresponding PyJWT verifications, blocking cross-tenant tokens
1066 minted by other applications that share the same IdP signing keys.
1067 When both are unset PyJWT only checks the signature and expiry, which
1068 is preserved for backward compatibility but logged once as a warning.
1070 The warning fires even in mixed deployments that also configure
1071 ``LiteLLM_JWTAuth.issuers``: tokens whose ``iss`` does not match any
1072 configured issuer fall through to this global path, and if env-var
1073 scoping is absent that fallback is itself unscoped.
1074 """
1075 audience: Final = os.getenv("JWT_AUDIENCE")
1076 issuer: Final = os.getenv("JWT_ISSUER")
1078 if audience is None and issuer is None and not cls._unscoped_jwt_warning_emitted:
1079 verbose_proxy_logger.warning(
1080 "JWT auth is enabled but neither JWT_AUDIENCE nor JWT_ISSUER "
1081 "is configured. Tokens minted by any application that shares "
1082 "the same IdP signing keys will be accepted. Set JWT_AUDIENCE "
1083 "(and ideally JWT_ISSUER) to scope this proxy."
1084 )
1085 cls._unscoped_jwt_warning_emitted = True
1087 options: Final[dict] = {}
1088 if audience is None:
1089 options["verify_aud"] = False
1090 if issuer is None:
1091 options["verify_iss"] = False
1093 return {
1094 "audience": audience,
1095 "issuer": issuer,
1096 "options": options or None,
1097 }
1099 def _get_configured_issuer(self, token: str) -> JWTIssuerConfig | None:
1100 litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
1101 if litellm_jwtauth is None:
1102 return None
1104 issuer_configs: Final = litellm_jwtauth.issuers
1105 if not issuer_configs:
1106 return None
1108 claims: Final = self.get_unverified_claims(token=token)
1109 if claims is None:
1110 return None
1112 issuer: Final = claims.get("iss")
1113 if not isinstance(issuer, str) or not issuer:
1114 return None
1116 for issuer_config in issuer_configs:
1117 if issuer_config.issuer == issuer:
1118 return issuer_config
1120 return None
1122 def _get_jwks_url_for_issuer(self, issuer_config: JWTIssuerConfig) -> str:
1123 if issuer_config.jwks_url:
1124 return issuer_config.jwks_url
1125 # _resolve_jwks_url fetches this OIDC discovery document and follows
1126 # its jwks_uri, matching JWTIssuerConfig.jwks_url's documented fallback.
1127 return f"{issuer_config.issuer.rstrip('/')}/.well-known/openid-configuration"
1129 def _get_claim_value_for_issuer_mapping(self, token: dict, claim_field: str) -> Any:
1130 """Resolve a mapped claim from ``token``.
1132 Returns ``None`` when the field is absent or empty so that mapped claims
1133 behave like the global ``litellm_jwtauth`` path — present claims override
1134 the normalised value, missing ones simply leave it ``None``.
1135 """
1136 sentinel: Final = object()
1137 claim_value: Final = get_nested_value(
1138 data=token,
1139 key_path=claim_field,
1140 default=sentinel,
1141 )
1142 if claim_value is sentinel or claim_value is None or claim_value == "":
1143 return None
1144 return claim_value
1146 def _apply_issuer_claim_mappings(self, token: dict, issuer_config: JWTIssuerConfig) -> dict:
1147 normalized: Final[dict] = {k: v for k, v in token.items() if k not in self.LITELLM_INTERNAL_CLAIMS}
1148 normalized[self.LITELLM_JWT_ISSUER_CLAIM] = issuer_config.issuer
1149 claim_mappings: Final = [
1150 (issuer_config.user_id_jwt_field, self.LITELLM_USER_ID_CLAIM),
1151 (issuer_config.user_email_jwt_field, self.LITELLM_USER_EMAIL_CLAIM),
1152 (issuer_config.team_id_jwt_field, self.LITELLM_TEAM_ID_CLAIM),
1153 (issuer_config.team_ids_jwt_field, self.LITELLM_TEAM_IDS_CLAIM),
1154 (issuer_config.org_id_jwt_field, self.LITELLM_ORG_ID_CLAIM),
1155 (issuer_config.end_user_id_jwt_field, self.LITELLM_END_USER_ID_CLAIM),
1156 ]
1158 for source_claim, normalized_claim in claim_mappings:
1159 if source_claim is None:
1160 continue
1161 claim_value = self._get_claim_value_for_issuer_mapping(
1162 token=token,
1163 claim_field=source_claim,
1164 )
1165 if claim_value is not None:
1166 normalized[normalized_claim] = claim_value
1168 return normalized
1170 def _get_jwk_from_public_key(self, public_key: dict) -> dict:
1171 jwk: Final = {}
1172 for key in ["kty", "kid", "n", "e", "x", "y", "crv"]:
1173 if key in public_key:
1174 jwk[key] = public_key[key]
1175 return jwk
1177 def _get_decode_options(
1178 self,
1179 audience: str | list[str] | None,
1180 issuer: str | None = None,
1181 disable_audience_validation: bool = False,
1182 ) -> dict | None:
1183 # Disabling audience verification must be an explicit choice — never
1184 # an implicit consequence of ``audience`` being None. Otherwise a
1185 # caller that accidentally constructs a config with ``audience=None``
1186 # (bypassing the model validator) would silently lose audience
1187 # validation. Require callers to opt in via
1188 # ``disable_audience_validation=True``.
1189 if audience is None and not disable_audience_validation:
1190 raise ValueError("audience must be provided unless disable_audience_validation=True")
1191 options: Final[dict] = {}
1192 if audience is None:
1193 options["verify_aud"] = False
1194 if issuer is None:
1195 options["verify_iss"] = False
1196 return options or None
1198 def _decode_jwt_with_public_key(
1199 self,
1200 token: str,
1201 public_key: dict | str,
1202 audience: str | list[str] | None,
1203 issuer: str | None = None,
1204 options: dict | None = None,
1205 disable_audience_validation: bool = False,
1206 ) -> dict:
1207 decode_options: Final = (
1208 options
1209 if options is not None
1210 else self._get_decode_options(
1211 audience=audience,
1212 issuer=issuer,
1213 disable_audience_validation=disable_audience_validation,
1214 )
1215 )
1217 if isinstance(public_key, dict):
1218 public_key_obj: Final = PyJWK.from_dict(self._get_jwk_from_public_key(public_key=public_key)).key
1219 return jwt.decode(
1220 token,
1221 public_key_obj,
1222 algorithms=self.SUPPORTED_JWT_ALGORITHMS,
1223 options=decode_options,
1224 audience=audience,
1225 issuer=issuer,
1226 leeway=self.leeway,
1227 )
1229 cert: Final = x509.load_pem_x509_certificate(public_key.encode(), default_backend())
1230 key: Final = cert.public_key().public_bytes(
1231 serialization.Encoding.PEM,
1232 serialization.PublicFormat.SubjectPublicKeyInfo,
1233 )
1234 return jwt.decode(
1235 token,
1236 key,
1237 algorithms=self.SUPPORTED_JWT_ALGORITHMS,
1238 audience=audience,
1239 issuer=issuer,
1240 options=decode_options,
1241 leeway=self.leeway,
1242 )
1244 async def _auth_jwt_with_issuer(self, token: str, issuer_config: JWTIssuerConfig, kid: str | None) -> dict:
1245 try:
1246 public_key: Final = await self._get_public_key_from_jwks_url(
1247 jwks_url=self._get_jwks_url_for_issuer(issuer_config=issuer_config),
1248 kid=kid,
1249 )
1250 except JWKSUnreachableError as e:
1251 raise jwks_unavailable_exception(e) from e
1253 try:
1254 payload: Final = self._decode_jwt_with_public_key(
1255 token=token,
1256 public_key=public_key,
1257 audience=issuer_config.audience,
1258 issuer=issuer_config.issuer,
1259 disable_audience_validation=issuer_config.disable_audience_validation,
1260 )
1261 except jwt.ExpiredSignatureError:
1262 raise ProxyException(
1263 message="Token Expired",
1264 type=ProxyErrorTypes.expired_key,
1265 param=None,
1266 code=status.HTTP_401_UNAUTHORIZED,
1267 )
1268 except Exception as e:
1269 raise Exception(f"Validation fails: {e}")
1271 return self._apply_issuer_claim_mappings(
1272 token=payload,
1273 issuer_config=issuer_config,
1274 )
1276 async def auth_jwt(self, token: str) -> dict:
1277 header: Final = jwt.get_unverified_header(token)
1279 verbose_proxy_logger.debug("header: %s", header)
1281 kid: Final = header.get("kid", None)
1283 issuer_config: Final = self._get_configured_issuer(token=token)
1284 if issuer_config is not None:
1285 return await self._auth_jwt_with_issuer(
1286 token=token,
1287 issuer_config=issuer_config,
1288 kid=kid,
1289 )
1291 decode_kwargs: Final = self._build_decode_kwargs()
1293 public_key: Final = await self.get_public_key(kid=kid)
1295 if public_key is not None:
1296 try:
1297 payload: Final = self._decode_jwt_with_public_key(
1298 token=token,
1299 public_key=public_key,
1300 audience=decode_kwargs["audience"],
1301 issuer=decode_kwargs["issuer"],
1302 options=decode_kwargs["options"],
1303 )
1304 return {k: v for k, v in payload.items() if k not in self.LITELLM_INTERNAL_CLAIMS}
1306 except jwt.ExpiredSignatureError:
1307 raise ProxyException(
1308 message="Token Expired",
1309 type=ProxyErrorTypes.expired_key,
1310 param=None,
1311 code=status.HTTP_401_UNAUTHORIZED,
1312 )
1313 except Exception as e:
1314 raise Exception(f"Validation fails: {e}")
1316 raise Exception("Invalid JWT Submitted")
1318 async def close(self):
1319 await self.http_handler.close()
1322class JWTAuthManager:
1323 """Manages JWT authentication and authorization operations"""
1325 @staticmethod
1326 def can_rbac_role_call_route(
1327 rbac_role: RBAC_ROLES,
1328 general_settings: dict,
1329 route: str,
1330 ) -> Literal[True]:
1331 """
1332 Checks if user is allowed to access the route, based on their role.
1333 """
1334 role_based_routes: Final = get_role_based_routes(rbac_role=rbac_role, general_settings=general_settings)
1336 if role_based_routes is None or route is None:
1337 return True
1339 is_allowed: Final = _allowed_routes_check(
1340 user_route=route,
1341 allowed_routes=role_based_routes,
1342 )
1344 if not is_allowed:
1345 raise HTTPException(
1346 status_code=403,
1347 detail=f"Role={rbac_role} not allowed to call route={route}. Allowed routes={role_based_routes}",
1348 )
1350 return True
1352 @staticmethod
1353 def can_rbac_role_call_model(
1354 rbac_role: RBAC_ROLES,
1355 general_settings: dict,
1356 model: str | None,
1357 ) -> Literal[True]:
1358 """
1359 Checks if user is allowed to access the model, based on their role.
1360 """
1361 role_based_models: Final = get_role_based_models(rbac_role=rbac_role, general_settings=general_settings)
1362 if role_based_models is None or model is None:
1363 return True
1365 if model not in role_based_models:
1366 internal_message: Final = (
1367 f"Role={rbac_role} not allowed to call model={model}. Allowed models={role_based_models}"
1368 )
1369 raise ModelAccessDeniedHTTPException(
1370 internal_message=internal_message,
1371 status_code=403,
1372 detail=model_access_denied_client_message(model=model),
1373 )
1375 return True
1377 @staticmethod
1378 def check_scope_based_access(
1379 scope_mappings: list[ScopeMapping],
1380 scopes: list[str],
1381 request_data: dict,
1382 general_settings: dict,
1383 ) -> None:
1384 """
1385 Check if scope allows access to the requested model
1386 """
1387 if not scope_mappings:
1388 return
1390 allowed_models: Final = []
1391 for sm in scope_mappings:
1392 if sm.scope in scopes and sm.models:
1393 allowed_models.extend(sm.models)
1395 requested_model: Final = request_data.get("model")
1397 if not requested_model:
1398 return
1400 if requested_model not in allowed_models:
1401 internal_message: Final = f"model={requested_model} not allowed. Allowed_models={allowed_models}"
1402 raise ModelAccessDeniedHTTPException(
1403 internal_message=internal_message,
1404 status_code=403,
1405 detail={"error": model_access_denied_client_message(model=requested_model)},
1406 )
1407 return
1409 @staticmethod
1410 async def check_rbac_role(
1411 jwt_handler: JWTHandler,
1412 jwt_valid_token: dict,
1413 general_settings: dict,
1414 request_data: dict,
1415 route: str,
1416 rbac_role: RBAC_ROLES | None,
1417 ) -> None:
1418 """Validate RBAC role and model access permissions"""
1419 if jwt_handler.litellm_jwtauth.enforce_rbac is True:
1420 if rbac_role is None:
1421 raise HTTPException(
1422 status_code=403,
1423 detail="Unmatched token passed in. enforce_rbac is set to True. Token must belong to a proxy admin, team, or user.",
1424 )
1425 JWTAuthManager.can_rbac_role_call_model(
1426 rbac_role=rbac_role,
1427 general_settings=general_settings,
1428 model=request_data.get("model"),
1429 )
1430 JWTAuthManager.can_rbac_role_call_route(
1431 rbac_role=rbac_role,
1432 general_settings=general_settings,
1433 route=route,
1434 )
1436 @staticmethod
1437 async def check_admin_access(
1438 jwt_handler: JWTHandler,
1439 scopes: list,
1440 route: str,
1441 user_id: str | None,
1442 org_id: str | None,
1443 api_key: str,
1444 jwt_valid_token: dict | None = None,
1445 user_email: str | None = None,
1446 agent_id: str | None = None,
1447 ) -> JWTAuthBuilderResult | None:
1448 """Check admin status and route access permissions"""
1449 if not jwt_handler.is_admin(scopes=scopes):
1450 return None
1452 is_allowed: Final = allowed_routes_check(
1453 user_role=LitellmUserRoles.PROXY_ADMIN,
1454 user_route=route,
1455 litellm_proxy_roles=jwt_handler.litellm_jwtauth,
1456 )
1457 if not is_allowed:
1458 allowed_routes: Final[list[Any]] = jwt_handler.litellm_jwtauth.admin_allowed_routes
1459 actual_routes: Final = get_actual_routes(allowed_routes=allowed_routes)
1460 raise Exception(f"Admin not allowed to access this route. Route={route}, Allowed Routes={actual_routes}")
1462 return JWTAuthBuilderResult(
1463 is_proxy_admin=True,
1464 team_object=None,
1465 user_object=None,
1466 end_user_object=None,
1467 org_object=None,
1468 token=api_key,
1469 team_id=None,
1470 user_id=user_id,
1471 user_email=user_email,
1472 end_user_id=None,
1473 org_id=org_id,
1474 team_membership=None,
1475 jwt_claims=jwt_valid_token or {},
1476 agent_id=agent_id,
1477 )
1479 @staticmethod
1480 def resolve_agent_id(
1481 jwt_handler: JWTHandler,
1482 jwt_valid_token: Mapping[str, object],
1483 agent_registry: AgentLookup,
1484 ) -> str | None:
1485 agent_claim: Final = jwt_handler.get_agent_claim(token=jwt_valid_token)
1486 if agent_claim is None:
1487 return None
1488 agent: Final = agent_registry.get_agent_by_id(agent_id=agent_claim) or agent_registry.get_agent_by_name(
1489 agent_name=agent_claim
1490 )
1491 if agent is None:
1492 raise HTTPException(
1493 status_code=status.HTTP_403_FORBIDDEN,
1494 detail=f"No registered agent matches JWT claim {jwt_handler.litellm_jwtauth.agent_id_jwt_field}={agent_claim}",
1495 )
1496 return agent.agent_id
1498 @staticmethod
1499 async def find_and_validate_specific_team_id(
1500 jwt_handler: JWTHandler,
1501 jwt_valid_token: dict,
1502 prisma_client: PrismaClient | None,
1503 user_api_key_cache: UserApiKeyCache,
1504 parent_otel_span: Span | None,
1505 proxy_logging_obj: ProxyLogging,
1506 team_id_upsert: bool | None = None,
1507 ) -> tuple[str | None, LiteLLM_TeamTable | None]:
1508 """Find and validate specific team ID from team_id_jwt_field or team_alias_jwt_field"""
1509 individual_team_id = jwt_handler.get_team_id(token=jwt_valid_token, default_value=None)
1510 team_alias: Final = jwt_handler.get_team_alias(token=jwt_valid_token, default_value=None)
1512 # `get_team_id` silently substitutes `team_id_default` for a missing
1513 # JWT team_id claim. When the token actually carries an alias claim,
1514 # that substitution would mask the alias-resolved team, so prefer
1515 # alias resolution. `get_all_jwt_team_ids` ignores `team_id_default`;
1516 # an empty result means no real JWT team_id claim is present.
1517 if (
1518 team_alias
1519 and individual_team_id is not None
1520 and not jwt_handler.get_all_jwt_team_ids(token=jwt_valid_token)
1521 ):
1522 individual_team_id = None
1524 team_object: LiteLLM_TeamTable | None = None
1526 if individual_team_id:
1527 try:
1528 team_object = await get_team_object(
1529 team_id=individual_team_id,
1530 prisma_client=prisma_client,
1531 user_api_key_cache=user_api_key_cache,
1532 parent_otel_span=parent_otel_span,
1533 proxy_logging_obj=proxy_logging_obj,
1534 team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert
1535 if team_id_upsert is None
1536 else team_id_upsert,
1537 )
1538 return individual_team_id, team_object
1539 except HTTPException as e:
1540 if e.status_code != 404 or not jwt_handler.litellm_jwtauth.team_claim_fallback:
1541 raise
1542 # Claim doesn't map to a known team — defer to fallback.
1543 verbose_proxy_logger.debug(
1544 "JWT team_id claim '%s' did not resolve to a team: %s",
1545 individual_team_id,
1546 e.detail,
1547 )
1548 return None, None
1550 if team_alias:
1551 verbose_proxy_logger.info("JWT Auth: Resolving team by alias: '%s'", team_alias)
1552 team_object = await get_team_object_by_alias(
1553 team_alias=team_alias,
1554 prisma_client=prisma_client,
1555 user_api_key_cache=user_api_key_cache,
1556 parent_otel_span=parent_otel_span,
1557 proxy_logging_obj=proxy_logging_obj,
1558 )
1559 if team_object:
1560 individual_team_id = team_object.team_id
1561 verbose_proxy_logger.info(
1562 "JWT Auth: Resolved team_alias='%s' to team_id='%s'", team_alias, individual_team_id
1563 )
1564 return individual_team_id, team_object
1566 # Check if team is required but not found
1567 if jwt_handler.is_required_team_id() is True:
1568 team_id_field: Final = jwt_handler.litellm_jwtauth.team_id_jwt_field
1569 team_alias_field: Final = jwt_handler.litellm_jwtauth.team_alias_jwt_field
1570 hint = ""
1571 if team_id_field:
1572 # "roles.0" — dot-notation numeric indexing is not supported
1573 if "." in team_id_field:
1574 parts: Final = team_id_field.rsplit(".", 1)
1575 if parts[-1].isdigit():
1576 base_field = parts[0]
1577 hint = (
1578 f" Hint: dot-notation array indexing (e.g. '{team_id_field}') is not "
1579 f"supported. Use '{base_field}' instead — LiteLLM automatically "
1580 f"uses the first element when the field value is a list."
1581 )
1582 # "roles[0]" — bracket-notation indexing is also not supported in get_nested_value
1583 elif "[" in team_id_field and team_id_field.endswith("]"):
1584 m: Final = re.match(r"^(\w+)\[(\d+)\]$", team_id_field)
1585 if m:
1586 base_field = m.group(1)
1587 hint = (
1588 f" Hint: array indexing (e.g. '{team_id_field}') is not supported "
1589 f"in team_id_jwt_field. Use '{base_field}' instead — LiteLLM "
1590 f"automatically uses the first element when the field value is a list."
1591 )
1592 raise Exception(
1593 f"No team found in token. Checked team_id field '{team_id_field}' and team_alias field '{team_alias_field}'.{hint}"
1594 )
1596 return individual_team_id, team_object
1598 @staticmethod
1599 def get_all_team_ids(jwt_handler: JWTHandler, jwt_valid_token: dict) -> set[str]:
1600 """Get combined team IDs from groups and individual team_id"""
1601 team_ids_from_groups: Final = jwt_handler.get_team_ids_from_jwt(token=jwt_valid_token)
1603 all_team_ids: Final = set(team_ids_from_groups)
1605 return all_team_ids
1607 @staticmethod
1608 def _team_has_passthrough_route_access(
1609 team_object: LiteLLM_TeamTable | None,
1610 route: str,
1611 request_method: str | None = None,
1612 team_allowed_routes: Collection[str] = (),
1613 ) -> bool:
1614 normalized_request_method: Final = request_method.upper() if isinstance(request_method, str) else None
1615 if not RouteChecks.is_auth_enforced_pass_through_route(
1616 route=route,
1617 method=normalized_request_method,
1618 ):
1619 return True
1621 if RouteChecks.jwt_team_routes_grant_pass_through(route=route, team_allowed_routes=team_allowed_routes):
1622 return True
1624 # JWT team selection is team-scoped; key metadata is not available here,
1625 # so beyond the JWT config grant above, only the selected team's metadata grants access.
1626 return RouteChecks.check_passthrough_route_access(
1627 route=route,
1628 user_api_key_dict=UserAPIKeyAuth(team_metadata=(team_object.metadata or {}) if team_object else {}),
1629 )
1631 @staticmethod
1632 def _raise_team_passthrough_route_denial(route: str) -> None:
1633 raise HTTPException(
1634 status_code=403,
1635 detail=(
1636 f"Team not allowed to access passthrough route {route}. "
1637 "Configure `allowed_passthrough_routes` on the team."
1638 ),
1639 )
1641 @staticmethod
1642 async def find_team_with_model_access(
1643 team_ids: set[str],
1644 requested_model: str | None,
1645 route: str,
1646 jwt_handler: JWTHandler,
1647 prisma_client: PrismaClient | None,
1648 user_api_key_cache: UserApiKeyCache,
1649 parent_otel_span: Span | None,
1650 proxy_logging_obj: ProxyLogging,
1651 request_method: str | None = None,
1652 ) -> tuple[str | None, LiteLLM_TeamTable | None]:
1653 """Find first team with access to the requested model"""
1654 from litellm.proxy.proxy_server import llm_router
1656 denied_auth_enforced_pass_through_route = False
1658 if not team_ids:
1659 if (
1660 jwt_handler.litellm_jwtauth.enforce_team_based_model_access
1661 and not jwt_handler.litellm_jwtauth.fallback_to_db_teams
1662 ):
1663 raise HTTPException(
1664 status_code=403,
1665 detail="No teams found in token. `enforce_team_based_model_access` is set to True. Token must belong to a team.",
1666 )
1667 return None, None
1669 any_claim_team_resolved = False
1670 for team_id in team_ids:
1671 try:
1672 team_object = await get_team_object(
1673 team_id=team_id,
1674 prisma_client=prisma_client,
1675 user_api_key_cache=user_api_key_cache,
1676 parent_otel_span=parent_otel_span,
1677 proxy_logging_obj=proxy_logging_obj,
1678 )
1680 if team_object is not None:
1681 any_claim_team_resolved = True
1683 if team_object and team_object.models is not None:
1684 team_models = team_object.models
1685 if isinstance(team_models, list) and (
1686 not requested_model
1687 or await can_team_access_model(
1688 model=requested_model,
1689 team_object=team_object,
1690 llm_router=llm_router,
1691 team_model_aliases=team_model_aliases(team_object),
1692 )
1693 ):
1694 is_allowed = allowed_routes_check(
1695 user_role=LitellmUserRoles.TEAM,
1696 user_route=route,
1697 litellm_proxy_roles=jwt_handler.litellm_jwtauth,
1698 )
1699 if is_allowed and not JWTAuthManager._team_has_passthrough_route_access(
1700 team_object=team_object,
1701 route=route,
1702 request_method=request_method,
1703 team_allowed_routes=jwt_handler.litellm_jwtauth.team_allowed_routes,
1704 ):
1705 is_allowed = False
1706 denied_auth_enforced_pass_through_route = True
1707 verbose_proxy_logger.debug(
1708 "JWT team route check: team_id=%s, route=%s, is_allowed=%s", team_id, route, is_allowed
1709 )
1710 if is_allowed:
1711 return team_id, team_object
1712 except Exception:
1713 continue
1715 if denied_auth_enforced_pass_through_route:
1716 JWTAuthManager._raise_team_passthrough_route_denial(route=route)
1718 if requested_model and (any_claim_team_resolved or not jwt_handler.litellm_jwtauth.team_claim_fallback):
1719 # Claim resolved but no model access, or fallback disabled — deny.
1720 raise HTTPException(
1721 status_code=403,
1722 detail=f"No team has access to the requested model: {requested_model}. Checked teams={team_ids}. Check `/models` to see all available models.",
1723 )
1725 # No claim team resolved and fallback enabled — defer to fallback.
1726 return None, None
1728 @staticmethod
1729 async def get_user_info(
1730 jwt_handler: JWTHandler,
1731 jwt_valid_token: dict,
1732 ) -> tuple[str | None, str | None, bool | None]:
1733 """Get user email and validation status"""
1734 user_email: Final = jwt_handler.get_user_email(token=jwt_valid_token, default_value=None)
1735 valid_user_email = None
1736 if jwt_handler.is_enforced_email_domain():
1737 valid_user_email = False if user_email is None else jwt_handler.is_allowed_domain(user_email=user_email)
1738 user_id: Final = jwt_handler.get_user_id(token=jwt_valid_token, default_value=user_email)
1739 return user_id, user_email, valid_user_email
1741 @staticmethod
1742 def _canonical_user_id_from_db(
1743 user_id: str | None,
1744 user_object: LiteLLM_UserTable | None,
1745 ) -> str | None:
1746 """Id used for spend / team-membership attribution.
1748 JWT claim (often email) is only a lookup key. If fuzzy match in
1749 ``get_user_object`` resolved a legacy row with a different ``user_id``,
1750 use that row's id; otherwise keep the claim. GH #26789.
1751 """
1752 return canonical_user_id(user_id=user_id, user_object=user_object)
1754 @staticmethod
1755 async def get_objects(
1756 user_id: str | None,
1757 user_email: str | None,
1758 org_id: str | None,
1759 end_user_id: str | None,
1760 team_id: str | None,
1761 valid_user_email: bool | None,
1762 jwt_handler: JWTHandler,
1763 prisma_client: PrismaClient | None,
1764 user_api_key_cache: UserApiKeyCache,
1765 parent_otel_span: Span | None,
1766 proxy_logging_obj: ProxyLogging,
1767 route: str,
1768 org_alias: str | None = None,
1769 user_id_upsert: bool | None = None,
1770 ) -> tuple[
1771 LiteLLM_UserTable | None,
1772 LiteLLM_OrganizationTable | None,
1773 LiteLLM_EndUserTable | None,
1774 LiteLLM_TeamMembership | None,
1775 str | None,
1776 ]:
1777 """Get user, org, end-user, and team-membership objects.
1779 Returns ``(..., effective_user_id)``: JWT claim unless fuzzy lookup
1780 matched a legacy row (GH #26789).
1781 """
1783 # Get org object - first try by ID, then by alias
1784 org_object: LiteLLM_OrganizationTable | None = None
1785 if org_id:
1786 org_object = (
1787 await get_org_object(
1788 org_id=org_id,
1789 prisma_client=prisma_client,
1790 user_api_key_cache=user_api_key_cache,
1791 parent_otel_span=parent_otel_span,
1792 proxy_logging_obj=proxy_logging_obj,
1793 )
1794 if org_id
1795 else None
1796 )
1797 elif org_alias:
1798 verbose_proxy_logger.info("JWT Auth: Resolving org by alias: '%s'", org_alias)
1799 org_object = await get_org_object_by_alias(
1800 org_alias=org_alias,
1801 prisma_client=prisma_client,
1802 user_api_key_cache=user_api_key_cache,
1803 parent_otel_span=parent_otel_span,
1804 proxy_logging_obj=proxy_logging_obj,
1805 )
1806 if org_object:
1807 verbose_proxy_logger.info(
1808 "JWT Auth: Resolved org_alias='%s' to org_id='%s'", org_alias, org_object.organization_id
1809 )
1811 # Check if email domain is allowed before attempting to get/create user
1812 if valid_user_email is False:
1813 raise ProxyException(
1814 message=f"Email domain not allowed. User email: {user_email}. Allowed domain: {jwt_handler.litellm_jwtauth.user_allowed_email_domain}",
1815 type=ProxyErrorTypes.auth_error,
1816 param="user_email",
1817 code=403,
1818 )
1820 user_object, team_membership_object, effective_user_id = await GrantResolver(
1821 prisma_client,
1822 user_api_key_cache,
1823 parent_otel_span=parent_otel_span,
1824 proxy_logging_obj=proxy_logging_obj,
1825 load_user=get_user_object,
1826 load_team=get_team_object,
1827 load_membership=get_team_membership,
1828 ).resolve_identity(
1829 UserLookup(
1830 user_id=user_id,
1831 user_email=user_email,
1832 sso_user_id=user_id,
1833 upsert=(
1834 jwt_handler.is_upsert_user_id(valid_user_email=valid_user_email)
1835 if user_id_upsert is None
1836 else user_id_upsert
1837 ),
1838 ),
1839 team_id=team_id,
1840 )
1842 end_user_object: LiteLLM_EndUserTable | None = None
1843 if end_user_id:
1844 end_user_object = (
1845 await get_end_user_object(
1846 end_user_id=end_user_id,
1847 prisma_client=prisma_client,
1848 user_api_key_cache=user_api_key_cache,
1849 parent_otel_span=parent_otel_span,
1850 proxy_logging_obj=proxy_logging_obj,
1851 route=route,
1852 )
1853 if end_user_id
1854 else None
1855 )
1857 return (
1858 user_object,
1859 org_object,
1860 end_user_object,
1861 team_membership_object,
1862 effective_user_id,
1863 )
1865 @staticmethod
1866 def validate_object_id(
1867 user_id: str | None,
1868 team_id: str | None,
1869 enforce_rbac: bool,
1870 is_proxy_admin: bool,
1871 ) -> Literal[True]:
1872 """If enforce_rbac is true, validate that a valid rbac id is returned for spend tracking"""
1873 if enforce_rbac and not is_proxy_admin and not user_id and not team_id:
1874 raise HTTPException(
1875 status_code=403,
1876 detail="No user or team id found in token. enforce_rbac is set to True. Token must belong to a proxy admin, team, or user.",
1877 )
1878 return True
1880 @staticmethod
1881 def _team_header_value(request_headers: Mapping[str, str] | None) -> str | None:
1882 if not request_headers:
1883 return None
1884 normalized_headers: Final = {k.lower(): v for k, v in request_headers.items()}
1885 return normalized_headers.get("x-litellm-team-id")
1887 @staticmethod
1888 def _raise_header_team_not_allowed(header_value: str, allowed_team_ids: set[str]) -> NoReturn:
1889 raise HTTPException(
1890 status_code=403,
1891 detail=(
1892 f"x-litellm-team-id '{header_value}' does not resolve to a team id or a unique team alias in your "
1893 f"JWT's allowed teams. Allowed team ids: {sorted(allowed_team_ids)}"
1894 ),
1895 )
1897 @staticmethod
1898 async def _team_id_by_alias(
1899 team_alias: str,
1900 prisma_client: PrismaClient | None,
1901 user_api_key_cache: UserApiKeyCache,
1902 parent_otel_span: Span | None,
1903 proxy_logging_obj: ProxyLogging,
1904 ) -> str | None:
1905 if prisma_client is None:
1906 return None
1907 try:
1908 team: Final = await get_team_object_by_alias(
1909 team_alias=team_alias,
1910 prisma_client=prisma_client,
1911 user_api_key_cache=user_api_key_cache,
1912 parent_otel_span=parent_otel_span,
1913 proxy_logging_obj=proxy_logging_obj,
1914 )
1915 except HTTPException as exc:
1916 if exc.status_code >= 500:
1917 raise
1918 return None
1919 return team.team_id
1921 @staticmethod
1922 async def resolve_team_from_header(
1923 request_headers: Mapping[str, str] | None,
1924 allowed_team_ids: set[str],
1925 fallback_to_db_teams: bool,
1926 prisma_client: PrismaClient | None,
1927 user_api_key_cache: UserApiKeyCache,
1928 parent_otel_span: Span | None,
1929 proxy_logging_obj: ProxyLogging,
1930 ) -> HeaderTeam | None:
1931 """
1932 The team named by x-litellm-team-id, which may carry a team id or a team
1933 alias. A value that is already an allowed team id never costs a lookup;
1934 under the DB fallback any other value is accepted provisionally, by id
1935 or alias, for the membership check auth_builder runs later. Under the
1936 DB fallback only a team row that is provably absent falls through to
1937 the alias lookup; a read that failed for any other reason keeps the
1938 membership denial the id path already gives.
1940 Raises:
1941 HTTPException: 403 when neither the value nor the team it aliases is
1942 an allowed team, or the DB fallback's membership denial when the
1943 value names no team at all; an alias several teams share resolves
1944 to no team and is denied like an unknown value; a 5xx from the
1945 alias lookup itself is re-raised rather than reported as a denial
1946 """
1947 header_value: Final = JWTAuthManager._team_header_value(request_headers)
1948 if not header_value:
1949 return None
1951 if header_value in allowed_team_ids:
1952 verbose_proxy_logger.debug("Using team_id from x-litellm-team-id header: %s", header_value)
1953 return HeaderTeam(header_value=header_value, team_id=header_value)
1955 if fallback_to_db_teams:
1956 try:
1957 await get_team_object(
1958 team_id=header_value,
1959 prisma_client=prisma_client,
1960 user_api_key_cache=user_api_key_cache,
1961 parent_otel_span=parent_otel_span,
1962 proxy_logging_obj=proxy_logging_obj,
1963 team_id_upsert=False,
1964 )
1965 except TeamNotFoundError:
1966 aliased_team_id: Final = await JWTAuthManager._team_id_by_alias(
1967 header_value, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj
1968 )
1969 if aliased_team_id is None:
1970 JWTAuthManager._raise_header_team_membership_denial(header_value)
1971 return HeaderTeam(header_value=header_value, team_id=aliased_team_id)
1972 except HTTPException:
1973 JWTAuthManager._raise_header_team_membership_denial(header_value)
1974 return HeaderTeam(header_value=header_value, team_id=header_value)
1976 team_id_by_alias: Final = await JWTAuthManager._team_id_by_alias(
1977 header_value, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj
1978 )
1979 if team_id_by_alias is None or team_id_by_alias not in allowed_team_ids:
1980 JWTAuthManager._raise_header_team_not_allowed(header_value, allowed_team_ids)
1981 verbose_proxy_logger.debug("Using team_id %s for x-litellm-team-id alias: %s", team_id_by_alias, header_value)
1982 return HeaderTeam(header_value=header_value, team_id=team_id_by_alias)
1984 @staticmethod
1985 async def map_user_to_teams(
1986 user_object: LiteLLM_UserTable | None,
1987 team_object: LiteLLM_TeamTable | None,
1988 ):
1989 """
1990 Map user to teams.
1991 - If user is not in team, add them to the team
1992 - If user is in team, do nothing
1993 """
1994 from litellm.proxy.management_endpoints.team_endpoints import team_member_add
1996 if not user_object:
1997 return
1999 if not team_object:
2000 return
2002 # check if user is in team
2003 for member in team_object.members_with_roles:
2004 if member.user_id and member.user_id == user_object.user_id:
2005 return
2007 data: Final = TeamMemberAddRequest(
2008 member=Member(
2009 user_id=user_object.user_id,
2010 role="user", # [TODO]: allow controlling role within team based on jwt token
2011 ),
2012 team_id=team_object.team_id,
2013 )
2014 # add user to team - make this non-blocking to avoid authentication failures
2015 try:
2016 await team_member_add(
2017 data=data,
2018 user_api_key_dict=UserAPIKeyAuth(
2019 user_role=LitellmUserRoles.PROXY_ADMIN
2020 ), # [TODO]: expose an internal service role, for better tracking
2021 )
2022 verbose_proxy_logger.debug(
2023 "Successfully added user %s to team %s", user_object.user_id, team_object.team_id
2024 )
2025 except ProxyException as e:
2026 if e.type == ProxyErrorTypes.team_member_already_in_team:
2027 verbose_proxy_logger.debug(
2028 "User %s is already a member of team %s", user_object.user_id, team_object.team_id
2029 )
2030 return
2031 else:
2032 raise e
2033 return
2035 @staticmethod
2036 async def sync_user_role_and_teams(
2037 jwt_handler: JWTHandler,
2038 jwt_valid_token: dict,
2039 user_object: LiteLLM_UserTable | None,
2040 prisma_client: PrismaClient | None,
2041 user_api_key_cache: UserApiKeyCache | None = None,
2042 ) -> None:
2043 """
2044 Sync user role and team memberships with JWT claims
2046 The goal of this method is to ensure:
2047 1. The user role on LiteLLM DB is in sync with the IDP provider role
2048 2. The user is a member of the teams specified in the JWT token
2050 This method is only called if sync_user_role_and_teams is set to True in the JWT config.
2051 """
2052 if not jwt_handler.litellm_jwtauth.sync_user_role_and_teams:
2053 return
2055 if user_object is None or prisma_client is None:
2056 return
2058 # Update user role
2059 new_role: Final = jwt_handler.map_jwt_role_to_litellm_role(jwt_valid_token)
2060 if new_role and user_object.user_role != new_role.value:
2061 await UserRepository(prisma_client).table.update(
2062 where={"user_id": user_object.user_id},
2063 data={"user_role": new_role.value},
2064 )
2065 user_object.user_role = new_role.value
2066 if user_api_key_cache is not None:
2067 await user_api_key_cache.async_set_cache(
2068 key=user_object.user_id,
2069 value=user_object,
2070 model_type=LiteLLM_UserTable,
2071 ttl=get_management_object_ttl(user_api_key_cache),
2072 )
2074 # Sync team memberships. With fallback_to_db_teams on, read both plural and
2075 # singular claim shapes so a singular-only IdP token (e.g. Okta/Auth0) is
2076 # not mistaken for claimless and left with stale DB memberships the fallback
2077 # could later attribute. With the flag off, keep the upstream plural-only
2078 # reconciliation so existing deployments are unchanged.
2079 jwt_team_ids: Final = set(
2080 jwt_handler.get_all_jwt_team_ids(jwt_valid_token)
2081 if jwt_handler.litellm_jwtauth.fallback_to_db_teams
2082 else jwt_handler.get_team_ids_from_jwt(jwt_valid_token)
2083 )
2084 existing_teams: Final = set(user_object.teams or [])
2085 teams_to_add: Final = jwt_team_ids - existing_teams
2086 preserve_db_teams_without_claims: Final = jwt_handler.litellm_jwtauth.fallback_to_db_teams and not jwt_team_ids
2087 teams_to_remove: Final = set() if preserve_db_teams_without_claims else existing_teams - jwt_team_ids
2088 if teams_to_add or teams_to_remove:
2089 from litellm.proxy.management_endpoints.scim.scim_v2 import (
2090 patch_team_membership,
2091 )
2093 await patch_team_membership(
2094 user_id=user_object.user_id,
2095 teams_ids_to_add_user_to=list(teams_to_add),
2096 teams_ids_to_remove_user_from=list(teams_to_remove),
2097 )
2098 user_object.teams = list(jwt_team_ids)
2099 if user_api_key_cache is not None:
2100 await user_api_key_cache.async_set_cache(
2101 key=user_object.user_id,
2102 value=user_object,
2103 model_type=LiteLLM_UserTable,
2104 ttl=get_management_object_ttl(user_api_key_cache),
2105 )
2106 return
2108 @staticmethod
2109 async def _attach_team_from_header_for_admin(
2110 admin_result: JWTAuthBuilderResult,
2111 route: str,
2112 request_headers: Mapping[str, str] | None,
2113 jwt_handler: JWTHandler,
2114 prisma_client: PrismaClient | None,
2115 user_api_key_cache: UserApiKeyCache,
2116 parent_otel_span: Span | None,
2117 proxy_logging_obj: ProxyLogging,
2118 team_id_upsert: bool | None = None,
2119 ) -> None:
2120 """Attach team context from x-litellm-team-id to an admin result.
2122 Only applies on LLM API routes so team TPM/RPM limits and attribution
2123 are enforced when admins act on behalf of a team. Admin management
2124 routes ignore the header to preserve pre-existing bypass behavior.
2125 """
2126 header_team_id: Final = request_headers.get("x-litellm-team-id") if request_headers else None
2127 if not header_team_id or not RouteChecks.is_llm_api_route(route=route):
2128 return
2129 try:
2130 team_object: Final = await get_team_object(
2131 team_id=header_team_id,
2132 prisma_client=prisma_client,
2133 user_api_key_cache=user_api_key_cache,
2134 parent_otel_span=parent_otel_span,
2135 proxy_logging_obj=proxy_logging_obj,
2136 team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert if team_id_upsert is None else team_id_upsert,
2137 )
2138 except Exception as e:
2139 # Fall back to pre-PR admin behavior: honor the admin's
2140 # authorization but skip team attribution/limits for this
2141 # request. Log so operators can find the misconfigured caller.
2142 verbose_proxy_logger.warning(
2143 "admin x-litellm-team-id=%r on route=%s could not be resolved (%s); "
2144 "proceeding with admin access, team context NOT attached.",
2145 header_team_id,
2146 route,
2147 e,
2148 )
2149 return
2150 admin_result["team_id"] = header_team_id
2151 admin_result["team_object"] = team_object
2153 @staticmethod
2154 async def _resolve_single_team_fallback(
2155 user_object: LiteLLM_UserTable | None,
2156 user_id: str | None,
2157 prisma_client: PrismaClient | None,
2158 user_api_key_cache: UserApiKeyCache,
2159 parent_otel_span: Span | None,
2160 proxy_logging_obj: ProxyLogging,
2161 team_id_upsert: bool | None,
2162 ) -> tuple:
2163 """
2164 If JWT did not resolve team_id, but the user belongs to exactly one team
2165 in LiteLLM, load that team (and membership when user_id is set) so that
2166 spend / metadata can be attributed correctly.
2168 Returns (team_id, team_object, team_membership_object).
2169 A team that cannot be loaded (HTTPException from get_team_object) is
2170 debug-logged and the tuple is (None, None, None), the same as the DB
2171 team fallback. A failed membership read propagates, so a database
2172 outage surfaces as the 503 the rest of auth answers with instead of
2173 serving the request with the member's limits dropped.
2174 """
2175 if user_object is None or not user_object.teams or len(user_object.teams) != 1:
2176 return None, None, None
2178 _tid: Final = user_object.teams[0]
2179 try:
2180 team_row: Final = await get_team_object(
2181 team_id=_tid,
2182 prisma_client=prisma_client,
2183 user_api_key_cache=user_api_key_cache,
2184 parent_otel_span=parent_otel_span,
2185 proxy_logging_obj=proxy_logging_obj,
2186 team_id_upsert=team_id_upsert,
2187 )
2188 except HTTPException:
2189 verbose_proxy_logger.debug(
2190 "JWT single-team fallback: team could not be loaded, skipping. team_id=%s",
2191 _tid,
2192 exc_info=True,
2193 )
2194 return None, None, None
2195 if team_row is None:
2196 return None, None, None
2198 if not user_id:
2199 return _tid, team_row, None
2201 team_membership: Final = await get_team_membership(
2202 user_id=user_id,
2203 team_id=_tid,
2204 prisma_client=prisma_client,
2205 user_api_key_cache=user_api_key_cache,
2206 parent_otel_span=parent_otel_span,
2207 proxy_logging_obj=proxy_logging_obj,
2208 )
2209 return _tid, team_row, team_membership
2211 @staticmethod
2212 async def _resolve_db_team_fallback(
2213 user_object: LiteLLM_UserTable | None,
2214 user_id: str | None,
2215 requested_model: str | None,
2216 route: str,
2217 jwt_handler: JWTHandler,
2218 enforce_team_based_model_access: bool,
2219 team_id_upsert: bool,
2220 prisma_client: PrismaClient | None,
2221 user_api_key_cache: UserApiKeyCache,
2222 parent_otel_span: Span | None,
2223 proxy_logging_obj: ProxyLogging,
2224 request_method: str | None = None,
2225 ) -> tuple[str | None, LiteLLM_TeamTable | None, LiteLLM_TeamMembership | None]:
2226 """
2227 Resolve a team for a user whose JWT carries no team claims by selecting
2228 the first DB team membership that loads successfully and, when a model is
2229 requested, can access that model — mirroring the per-team model-access
2230 check the claim-based path enforces, so a team's `models` restriction is
2231 not bypassed by the fallback.
2233 The same `team_allowed_routes` gate the claim-based path applies is
2234 enforced here too, so a DB-selected team cannot reach a route the JWT
2235 config excludes for team-role callers. Auth-enforced passthrough routes
2236 are exempt from that gate by design (they are governed by the team's
2237 `allowed_passthrough_routes`, re-checked by the caller).
2239 The resolved team's membership row is loaded too (when user_id is set) so
2240 per-team membership budget limits are enforced on the fallback path the
2241 same as on the claim-based path.
2243 Raises HTTP 403 when the user has no usable DB team membership and
2244 `enforce_team_based_model_access` is set; otherwise returns (None, None, None).
2245 """
2246 from litellm.proxy.proxy_server import llm_router
2248 user_team_ids: Final = user_object.teams if user_object else []
2249 team_route_allowed: Final = JWTAuthManager._is_team_route_allowed(
2250 route=route, request_method=request_method, jwt_handler=jwt_handler
2251 )
2252 any_team_resolved = False
2253 for candidate_team_id in user_team_ids:
2254 try:
2255 team_object = await get_team_object(
2256 team_id=candidate_team_id,
2257 prisma_client=prisma_client,
2258 user_api_key_cache=user_api_key_cache,
2259 parent_otel_span=parent_otel_span,
2260 proxy_logging_obj=proxy_logging_obj,
2261 team_id_upsert=team_id_upsert,
2262 )
2263 except HTTPException:
2264 continue
2265 any_team_resolved = True
2266 if requested_model:
2267 try:
2268 await can_team_access_model(
2269 model=requested_model,
2270 team_object=team_object,
2271 llm_router=llm_router,
2272 team_model_aliases=team_model_aliases(team_object),
2273 )
2274 except ProxyException:
2275 continue
2276 if not team_route_allowed:
2277 continue
2278 verbose_proxy_logger.debug(
2279 "JWT DB team fallback: resolved team_id=%s from user DB membership",
2280 candidate_team_id,
2281 )
2282 if user_id:
2283 return (
2284 candidate_team_id,
2285 team_object,
2286 await get_team_membership(
2287 user_id=user_id,
2288 team_id=candidate_team_id,
2289 prisma_client=prisma_client,
2290 user_api_key_cache=user_api_key_cache,
2291 parent_otel_span=parent_otel_span,
2292 proxy_logging_obj=proxy_logging_obj,
2293 ),
2294 )
2295 return candidate_team_id, team_object, None
2297 if enforce_team_based_model_access:
2298 if requested_model and any_team_resolved:
2299 raise HTTPException(
2300 status_code=403,
2301 detail=(
2302 f"No team you are a member of has access to the requested "
2303 f"model: {requested_model}. Check `/models` to see the models "
2304 f"available to you."
2305 ),
2306 )
2307 raise HTTPException(
2308 status_code=403,
2309 detail=("User is not a member of any team. Add the user to a team via the LiteLLM UI or API."),
2310 )
2311 return None, None, None
2313 @staticmethod
2314 def _is_team_route_allowed(
2315 route: str,
2316 request_method: str | None,
2317 jwt_handler: JWTHandler,
2318 ) -> bool:
2319 """
2320 Whether a team-role caller may reach `route` per the JWT config's
2321 `team_allowed_routes`. Auth-enforced passthrough routes are exempt
2322 here; their team's `allowed_passthrough_routes` gate runs separately.
2323 """
2324 normalized_method: Final = request_method.upper() if isinstance(request_method, str) else None
2325 return RouteChecks.is_auth_enforced_pass_through_route(
2326 route=route, method=normalized_method
2327 ) or allowed_routes_check(
2328 user_role=LitellmUserRoles.TEAM,
2329 user_route=route,
2330 litellm_proxy_roles=jwt_handler.litellm_jwtauth,
2331 )
2333 @staticmethod
2334 def _raise_header_team_membership_denial(header_value: str) -> NoReturn:
2335 """
2336 The single denial shape for a provisional x-litellm-team-id header,
2337 raised identically for nonexistent teams and for teams the user is not
2338 a member of, and naming only the value the caller sent, so the response
2339 reveals neither whether a team exists nor which id an alias maps to.
2340 """
2341 raise HTTPException(
2342 status_code=403,
2343 detail=(
2344 f"x-litellm-team-id '{header_value}' does not resolve to a team id or a unique team alias among your "
2345 "team memberships."
2346 ),
2347 )
2349 @staticmethod
2350 def _validate_header_team_in_db_membership(
2351 team_id: str,
2352 user_object: LiteLLM_UserTable | None,
2353 header_value: str,
2354 ) -> None:
2355 """
2356 A provisional team_id from the x-litellm-team-id header (accepted under
2357 fallback_to_db_teams because it is outside the JWT's teams) must exist in
2358 the user's DB team memberships before it becomes request context. The denial
2359 names `header_value`, the id or alias the caller sent, not `team_id`.
2360 """
2361 user_team_ids: Final = user_object.teams if user_object else []
2362 if team_id in user_team_ids:
2363 return
2364 JWTAuthManager._raise_header_team_membership_denial(header_value)
2366 @staticmethod
2367 async def auth_builder(
2368 api_key: str,
2369 jwt_handler: JWTHandler,
2370 request_data: dict,
2371 general_settings: dict,
2372 route: str,
2373 prisma_client: PrismaClient | None,
2374 user_api_key_cache: UserApiKeyCache,
2375 parent_otel_span: Span | None,
2376 proxy_logging_obj: ProxyLogging,
2377 request_headers: Mapping[str, str] | None = None,
2378 request_method: str | None = None,
2379 ) -> JWTAuthBuilderResult:
2380 return await JWTAuthManager.authorize_jwt(
2381 api_key=api_key,
2382 jwt_handler=jwt_handler,
2383 request_data=request_data,
2384 general_settings=general_settings,
2385 route=route,
2386 prisma_client=prisma_client,
2387 user_api_key_cache=user_api_key_cache,
2388 parent_otel_span=parent_otel_span,
2389 proxy_logging_obj=proxy_logging_obj,
2390 request_headers=request_headers,
2391 request_method=request_method,
2392 provisioning=_JWTProvisioning(
2393 user_id_upsert=jwt_handler.litellm_jwtauth.user_id_upsert,
2394 team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
2395 ),
2396 )
2398 @staticmethod
2399 async def authenticate_jwt(api_key: str, jwt_handler: JWTHandler) -> dict[str, object]:
2400 claims: Final = (
2401 await jwt_handler.get_oidc_userinfo(token=api_key)
2402 if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not jwt_handler.is_jwt(token=api_key)
2403 else await jwt_handler.auth_jwt(token=api_key)
2404 )
2405 validate: Final = jwt_handler.litellm_jwtauth.custom_validate
2406 if validate is not None and not validate(claims):
2407 raise HTTPException(status_code=403, detail="Invalid JWT token")
2408 return claims
2410 @staticmethod
2411 async def resolve_identity(
2412 api_key: str,
2413 jwt_handler: JWTHandler,
2414 prisma_client: PrismaClient | None,
2415 user_api_key_cache: UserApiKeyCache,
2416 parent_otel_span: Span | None,
2417 proxy_logging_obj: ProxyLogging,
2418 ) -> JWTIdentity:
2419 claims: Final = await JWTAuthManager.authenticate_jwt(api_key, jwt_handler)
2420 return await JWTAuthManager._resolve_claim_identity(
2421 claims, jwt_handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj
2422 )
2424 @staticmethod
2425 async def _resolve_claim_identity(
2426 claims: dict[str, object],
2427 jwt_handler: JWTHandler,
2428 prisma_client: PrismaClient | None,
2429 user_api_key_cache: UserApiKeyCache,
2430 parent_otel_span: Span | None,
2431 proxy_logging_obj: ProxyLogging,
2432 ) -> JWTIdentity:
2433 claim_user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(jwt_handler, claims)
2434 user_id: Final = (
2435 jwt_handler.get_object_id(token=claims, default_value=None) or claim_user_id
2436 if jwt_handler.get_rbac_role(token=claims) == LitellmUserRoles.INTERNAL_USER
2437 else claim_user_id
2438 )
2439 agent_id: Final = JWTAuthManager.resolve_agent_id(jwt_handler, claims, jwt_handler.agent_lookup)
2440 is_admin: Final = jwt_handler.is_admin(scopes=jwt_handler.get_scopes(token=claims))
2441 try:
2442 user, _, _, _, canonical_id = await JWTAuthManager.get_objects(
2443 user_id=user_id,
2444 user_email=user_email,
2445 org_id=None,
2446 end_user_id=None,
2447 team_id=None,
2448 valid_user_email=valid_user_email,
2449 jwt_handler=jwt_handler,
2450 prisma_client=prisma_client,
2451 user_api_key_cache=user_api_key_cache,
2452 parent_otel_span=parent_otel_span,
2453 proxy_logging_obj=proxy_logging_obj,
2454 route="",
2455 user_id_upsert=False,
2456 )
2457 except UserNotFoundError:
2458 if not is_admin:
2459 raise
2460 return JWTIdentity(user_id=user_id, user_object=None, agent_id=agent_id)
2461 return JWTIdentity(user_id=user_id if is_admin else canonical_id, user_object=user, agent_id=agent_id)
2463 @staticmethod
2464 async def authorize_jwt(
2465 api_key: str,
2466 jwt_handler: JWTHandler,
2467 request_data: dict[str, object],
2468 general_settings: dict[str, object],
2469 route: str,
2470 prisma_client: PrismaClient | None,
2471 user_api_key_cache: UserApiKeyCache,
2472 parent_otel_span: Span | None,
2473 proxy_logging_obj: ProxyLogging,
2474 request_headers: Mapping[str, str] | None = None,
2475 request_method: str | None = None,
2476 provisioning: _JWTProvisioning | None = None,
2477 ) -> JWTAuthBuilderResult:
2478 """Resolve and authorize JWT context; only normal admission supplies provisioning."""
2479 handler: Final = jwt_handler
2480 jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler)
2481 team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False
2482 model: Final = request_data.get("model")
2483 requested_model: Final = model if isinstance(model, str) else None
2485 # Check RBAC
2486 rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token)
2487 await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role)
2489 # Check Scope Based Access
2490 scopes: Final = handler.get_scopes(token=jwt_valid_token)
2491 if handler.litellm_jwtauth.enforce_scope_based_access and handler.litellm_jwtauth.scope_mappings:
2492 JWTAuthManager.check_scope_based_access(
2493 scope_mappings=handler.litellm_jwtauth.scope_mappings,
2494 scopes=scopes,
2495 request_data=request_data,
2496 general_settings=general_settings,
2497 )
2499 object_id = handler.get_object_id(token=jwt_valid_token, default_value=None)
2501 # Get basic user info
2502 user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token)
2504 # Get IDs
2505 org_id: Final = handler.get_org_id(token=jwt_valid_token, default_value=None)
2506 end_user_id: Final = handler.get_end_user_id(token=jwt_valid_token, default_value=None)
2507 team_id: str | None = None
2508 team_object: LiteLLM_TeamTable | None = None
2509 object_id = handler.get_object_id(token=jwt_valid_token, default_value=None)
2511 if rbac_role and object_id:
2512 if rbac_role == LitellmUserRoles.TEAM:
2513 team_id = object_id
2514 elif rbac_role == LitellmUserRoles.INTERNAL_USER:
2515 user_id = object_id
2517 agent_id: Final = JWTAuthManager.resolve_agent_id(
2518 jwt_handler=handler,
2519 jwt_valid_token=jwt_valid_token,
2520 agent_registry=handler.agent_lookup,
2521 )
2523 # Check admin access
2524 admin_result: Final = await JWTAuthManager.check_admin_access(
2525 handler,
2526 scopes,
2527 route,
2528 user_id,
2529 org_id,
2530 api_key,
2531 jwt_valid_token,
2532 user_email=user_email,
2533 agent_id=agent_id,
2534 )
2535 if admin_result:
2536 await JWTAuthManager._attach_team_from_header_for_admin(
2537 admin_result=admin_result,
2538 route=route,
2539 request_headers=request_headers,
2540 jwt_handler=handler,
2541 prisma_client=prisma_client,
2542 user_api_key_cache=user_api_key_cache,
2543 parent_otel_span=parent_otel_span,
2544 proxy_logging_obj=proxy_logging_obj,
2545 team_id_upsert=team_id_upsert,
2546 )
2547 if provisioning is None:
2548 identity: Final = await JWTAuthManager._resolve_claim_identity(
2549 jwt_valid_token, handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj
2550 )
2551 return {**admin_result, "user_object": identity.user_object}
2552 if prisma_client is None:
2553 return admin_result
2554 try:
2555 admin_user: Final = await get_user_object(
2556 user_id=user_id,
2557 user_email=user_email,
2558 sso_user_id=user_id,
2559 prisma_client=prisma_client,
2560 user_api_key_cache=user_api_key_cache,
2561 user_id_upsert=False,
2562 parent_otel_span=parent_otel_span,
2563 proxy_logging_obj=proxy_logging_obj,
2564 )
2565 except UserNotFoundError:
2566 return admin_result
2567 return {**admin_result, "user_object": admin_user}
2569 # Get team with model access
2570 ## Check if team_id is specified via x-litellm-team-id header
2571 all_team_ids: Final = JWTAuthManager.get_all_team_ids(handler, jwt_valid_token)
2572 specific_team_id: Final = handler.get_team_id(token=jwt_valid_token, default_value=None)
2574 # The DB fallback only applies when the token carries no team identity at
2575 # all. `get_all_jwt_team_ids` ignores `team_id_default` so a configured
2576 # default does not hide a claimless token, `get_team_alias` covers
2577 # alias-only tokens so the alias still resolves via
2578 # `find_and_validate_specific_team_id`, and `team_id is None` excludes
2579 # the RBAC team-role path (which already set `team_id`); otherwise a
2580 # provisional x-litellm-team-id header could override an RBAC-asserted team.
2581 db_team_fallback: Final = (
2582 handler.litellm_jwtauth.fallback_to_db_teams
2583 and not handler.get_all_jwt_team_ids(token=jwt_valid_token)
2584 and not handler.get_team_alias(token=jwt_valid_token, default_value=None)
2585 and team_id is None
2586 )
2587 if specific_team_id and not db_team_fallback:
2588 all_team_ids.add(specific_team_id)
2590 header_db_fallback: Final = handler.litellm_jwtauth.fallback_to_db_teams and team_id is None
2592 header_team: Final = await JWTAuthManager.resolve_team_from_header(
2593 request_headers=request_headers,
2594 allowed_team_ids=all_team_ids,
2595 fallback_to_db_teams=header_db_fallback,
2596 prisma_client=prisma_client,
2597 user_api_key_cache=user_api_key_cache,
2598 parent_otel_span=parent_otel_span,
2599 proxy_logging_obj=proxy_logging_obj,
2600 )
2601 provisional_header_team: Final = (
2602 header_team
2603 if header_team is not None and header_db_fallback and header_team.team_id not in all_team_ids
2604 else None
2605 )
2606 if header_team:
2607 team_id = header_team.team_id
2608 # A provisional header team (accepted because it is outside the
2609 # JWT's teams under fallback_to_db_teams) is validated against DB
2610 # membership further down; never upsert it here or an
2611 # attacker-supplied x-litellm-team-id would create an orphaned team
2612 # row before that check runs. A genuine membership team already
2613 # exists, so suppressing the upsert in that case costs nothing.
2614 try:
2615 team_object = await get_team_object(
2616 team_id=team_id,
2617 prisma_client=prisma_client,
2618 user_api_key_cache=user_api_key_cache,
2619 parent_otel_span=parent_otel_span,
2620 proxy_logging_obj=proxy_logging_obj,
2621 team_id_upsert=(team_id_upsert and provisional_header_team is None),
2622 )
2623 except HTTPException:
2624 if provisional_header_team is None:
2625 raise
2626 JWTAuthManager._raise_header_team_membership_denial(header_team.header_value)
2627 elif not team_id and not db_team_fallback:
2628 ## SPECIFIC TEAM ID
2629 (
2630 team_id,
2631 team_object,
2632 ) = await JWTAuthManager.find_and_validate_specific_team_id(
2633 handler,
2634 jwt_valid_token,
2635 prisma_client,
2636 user_api_key_cache,
2637 parent_otel_span,
2638 proxy_logging_obj,
2639 team_id_upsert=team_id_upsert,
2640 )
2642 if not team_object and not team_id:
2643 ## CHECK USER GROUP ACCESS
2644 team_id, team_object = await JWTAuthManager.find_team_with_model_access(
2645 team_ids=all_team_ids,
2646 requested_model=requested_model,
2647 route=route,
2648 request_method=request_method,
2649 jwt_handler=handler,
2650 prisma_client=prisma_client,
2651 user_api_key_cache=user_api_key_cache,
2652 parent_otel_span=parent_otel_span,
2653 proxy_logging_obj=proxy_logging_obj,
2654 )
2656 # The RBAC role-claim path (rbac_role == TEAM) sets team_id without
2657 # loading team_object, so fetch it here before gating an auth-enforced
2658 # passthrough route on the team's allowed_passthrough_routes.
2659 if (
2660 team_id
2661 and team_object is None
2662 and RouteChecks.is_auth_enforced_pass_through_route(
2663 route=route,
2664 method=(request_method.upper() if isinstance(request_method, str) else None),
2665 )
2666 ):
2667 team_object = await get_team_object(
2668 team_id=team_id,
2669 prisma_client=prisma_client,
2670 user_api_key_cache=user_api_key_cache,
2671 parent_otel_span=parent_otel_span,
2672 proxy_logging_obj=proxy_logging_obj,
2673 team_id_upsert=team_id_upsert,
2674 )
2676 if team_id and not JWTAuthManager._team_has_passthrough_route_access(
2677 team_object=team_object,
2678 route=route,
2679 request_method=request_method,
2680 team_allowed_routes=handler.litellm_jwtauth.team_allowed_routes,
2681 ):
2682 JWTAuthManager._raise_team_passthrough_route_denial(route=route)
2684 # Extract alias fields for resolution (if configured)
2685 org_alias: Final = handler.get_org_alias(token=jwt_valid_token, default_value=None)
2687 # get_objects returns effective_user_id for downstream spend attribution (GH #26789).
2688 (
2689 user_object,
2690 org_object,
2691 end_user_object,
2692 team_membership_object,
2693 user_id,
2694 ) = await JWTAuthManager.get_objects(
2695 user_id=user_id,
2696 user_email=user_email,
2697 org_id=org_id,
2698 end_user_id=end_user_id,
2699 team_id=team_id,
2700 valid_user_email=valid_user_email,
2701 jwt_handler=handler,
2702 prisma_client=prisma_client,
2703 user_api_key_cache=user_api_key_cache,
2704 parent_otel_span=parent_otel_span,
2705 proxy_logging_obj=proxy_logging_obj,
2706 route=route,
2707 org_alias=org_alias,
2708 user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False,
2709 )
2711 # Derive org_id from org_object if resolved by alias
2712 resolved_org_id: Final = org_object.organization_id if org_object else org_id
2714 if provisioning is not None:
2715 await JWTAuthManager.sync_user_role_and_teams(
2716 jwt_handler=handler,
2717 jwt_valid_token=jwt_valid_token,
2718 user_object=user_object,
2719 prisma_client=prisma_client,
2720 user_api_key_cache=user_api_key_cache,
2721 )
2723 # If JWT did not resolve team_id, attempt a team fallback.
2724 if team_id is None and db_team_fallback:
2725 (
2726 team_id,
2727 team_object,
2728 team_membership_object,
2729 ) = await JWTAuthManager._resolve_db_team_fallback(
2730 user_object=user_object,
2731 user_id=user_id,
2732 requested_model=requested_model,
2733 route=route,
2734 jwt_handler=handler,
2735 enforce_team_based_model_access=handler.litellm_jwtauth.enforce_team_based_model_access,
2736 team_id_upsert=team_id_upsert,
2737 prisma_client=prisma_client,
2738 user_api_key_cache=user_api_key_cache,
2739 parent_otel_span=parent_otel_span,
2740 proxy_logging_obj=proxy_logging_obj,
2741 request_method=request_method,
2742 )
2743 # The earlier passthrough gate ran when team_id was None; re-check
2744 # against the DB-resolved team so a fallback-selected team must also
2745 # pass the auth-enforced passthrough allowlist.
2746 if team_id and not JWTAuthManager._team_has_passthrough_route_access(
2747 team_object=team_object,
2748 route=route,
2749 request_method=request_method,
2750 team_allowed_routes=handler.litellm_jwtauth.team_allowed_routes,
2751 ):
2752 JWTAuthManager._raise_team_passthrough_route_denial(route=route)
2753 elif team_id is None:
2754 (
2755 team_id,
2756 team_object,
2757 team_membership_object,
2758 ) = await JWTAuthManager._resolve_single_team_fallback(
2759 user_object=user_object,
2760 user_id=user_id,
2761 prisma_client=prisma_client,
2762 user_api_key_cache=user_api_key_cache,
2763 parent_otel_span=parent_otel_span,
2764 proxy_logging_obj=proxy_logging_obj,
2765 team_id_upsert=team_id_upsert,
2766 )
2767 elif provisional_header_team is not None and team_id == provisional_header_team.team_id:
2768 JWTAuthManager._validate_header_team_in_db_membership(
2769 team_id=team_id,
2770 user_object=user_object,
2771 header_value=provisional_header_team.header_value,
2772 )
2773 if not JWTAuthManager._is_team_route_allowed(
2774 route=route,
2775 request_method=request_method,
2776 jwt_handler=handler,
2777 ):
2778 raise HTTPException(
2779 status_code=403,
2780 detail=(
2781 f"Team '{provisional_header_team.header_value}' (from x-litellm-team-id header) "
2782 f"is not allowed to access route '{route}'."
2783 ),
2784 )
2786 ## MAP USER TO TEAMS
2787 if provisioning is not None:
2788 await JWTAuthManager.map_user_to_teams(
2789 user_object=user_object,
2790 team_object=team_object,
2791 )
2793 # Validate that a valid rbac id is returned for spend tracking
2794 JWTAuthManager.validate_object_id(
2795 user_id=user_id,
2796 team_id=team_id,
2797 enforce_rbac=bool(general_settings.get("enforce_rbac", False)),
2798 is_proxy_admin=False,
2799 )
2801 # check if user is proxy admin
2802 is_proxy_admin: Final = bool(user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN)
2804 return JWTAuthBuilderResult(
2805 is_proxy_admin=is_proxy_admin,
2806 team_id=team_id,
2807 team_object=team_object,
2808 user_id=user_id,
2809 user_email=(user_object.user_email if user_object is not None and user_object.user_email else user_email),
2810 user_object=user_object,
2811 org_id=resolved_org_id, # Use resolved org_id (from alias lookup if applicable)
2812 org_object=org_object,
2813 end_user_id=end_user_id,
2814 end_user_object=end_user_object,
2815 token=api_key,
2816 team_membership=team_membership_object,
2817 jwt_claims=jwt_valid_token,
2818 agent_id=agent_id,
2819 )
2821 @staticmethod
2822 def user_api_key_auth_from_result(
2823 result: JWTAuthBuilderResult,
2824 parent_otel_span: Span | None = None,
2825 ) -> UserAPIKeyAuth:
2826 """Keep JWT identity and permission attribution identical across consumers."""
2827 user: Final = result["user_object"]
2828 admin: Final = result["is_proxy_admin"]
2829 return UserAPIKeyAuth(
2830 api_key=None,
2831 user_role=(
2832 LitellmUserRoles.PROXY_ADMIN
2833 if admin
2834 else LitellmUserRoles(user.user_role)
2835 if user is not None and user.user_role is not None
2836 else LitellmUserRoles.INTERNAL_USER
2837 ),
2838 user_id=result["user_id"],
2839 user_email=result["user_email"],
2840 team_id=result["team_id"],
2841 org_id=result["org_id"],
2842 end_user_id=result["end_user_id"],
2843 parent_otel_span=parent_otel_span,
2844 jwt_claims=result["jwt_claims"],
2845 agent_id=result.get("agent_id"),
2846 user_tpm_limit=user.tpm_limit if user is not None and not admin else None,
2847 user_rpm_limit=user.rpm_limit if user is not None and not admin else None,
2848 user_model_max_budget=user.model_max_budget if user is not None and not admin else None,
2849 **team_grants(
2850 team_object=result["team_object"],
2851 team_membership=result.get("team_membership"),
2852 user_id=result["user_id"],
2853 ),
2854 )