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

1""" 

2Supports using JWT's for authenticating into the proxy. 

3 

4Currently only supports admin. 

5 

6JWT token must have 'litellm_proxy_admin' in scope. 

7""" 

8 

9from __future__ import annotations 

10 

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 

20 

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 

29 

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 

71 

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) 

87 

88 

89class NoMatchingJWTPublicKeyError(Exception): 

90 """Raised when a JWKS endpoint returns no key matching the requested ``kid``.""" 

91 

92 

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.""" 

95 

96 

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_" 

103 

104_CachedValueT = TypeVar("_CachedValueT", bound=JWKKeyValue | str) 

105 

106 

107class _JWTAuthSettings(Protocol): 

108 """The JWT auth settings block this handler reads back through ``getattr``, when one is configured.""" 

109 

110 @property 

111 def issuers(self) -> Sequence[JWTIssuerConfig] | None: ... 111 ↛ exitline 111 didn't return from function 'issuers' because

112 

113 @property 

114 def public_key_ttl(self) -> float: ... 114 ↛ exitline 114 didn't return from function 'public_key_ttl' because

115 

116 @property 

117 def public_key_stale_ttl(self) -> float: ... 117 ↛ exitline 117 didn't return from function 'public_key_stale_ttl' because

118 

119 

120class _OIDCDiscoveryBody(TypedDict, total=False): 

121 """Decoded OIDC discovery document, read for the JWKS endpoint it advertises.""" 

122 

123 jwks_uri: ReadOnly[str] 

124 

125 

126class _OIDCDiscoveryResponse(Protocol): 

127 """The discovery endpoint's HTTP response, read for the decoded document it carries.""" 

128 

129 def json(self) -> _OIDCDiscoveryBody: ... 129 ↛ exitline 129 didn't return from function 'json' because

130 

131 

132class _UserInfoResponse(Protocol): 

133 """The OIDC UserInfo endpoint's HTTP response, read for the identity document it carries.""" 

134 

135 def json(self) -> dict[str, object]: ... 135 ↛ exitline 135 didn't return from function 'json' because

136 

137 

138@dataclass(frozen=True, slots=True) 

139class JWTIdentity: 

140 user_id: str | None 

141 user_object: LiteLLM_UserTable | None 

142 agent_id: str | None 

143 

144 

145@dataclass(frozen=True, slots=True) 

146class _JWTProvisioning: 

147 user_id_upsert: bool 

148 team_id_upsert: bool 

149 

150 

151@dataclass(frozen=True, slots=True) 

152class HeaderTeam: 

153 header_value: str 

154 team_id: str 

155 

156 

157class AgentLookup(Protocol): 

158 """The registered-agent lookups a JWT agent claim is matched against.""" 

159 

160 def get_agent_by_id(self, agent_id: str) -> AgentResponse | None: 

161 """The agent registered under ``agent_id``, if any.""" 

162 

163 def get_agent_by_name(self, agent_name: str) -> AgentResponse | None: 

164 """The agent registered under ``agent_name``, if any.""" 

165 

166 

167class _NoRegisteredAgents: 

168 """The lookup in force until the proxy binds its agent registry: no agent is registered, so no claim matches.""" 

169 

170 def get_agent_by_id(self, agent_id: str) -> None: 

171 return None 

172 

173 def get_agent_by_name(self, agent_name: str) -> None: 

174 return None 

175 

176 

177def _discovery_document(response: _OIDCDiscoveryResponse) -> _OIDCDiscoveryBody: 

178 """Decode an OIDC discovery response body.""" 

179 return response.json() 

180 

181 

182def _userinfo_document(response: _UserInfoResponse) -> dict[str, object]: 

183 """Decode an OIDC UserInfo response body into its JSON object form.""" 

184 return response.json() 

185 

186 

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 ) 

197 

198 

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

206 

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 ) 

240 

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

249 

250 def bind_agent_lookup(self, agent_lookup: AgentLookup) -> None: 

251 self.agent_lookup = agent_lookup 

252 

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 

264 

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 

271 

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 

280 

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 

293 

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. 

297 

298 Args: 

299 token (dict): The JWT token containing role information 

300 

301 Returns: 

302 Optional[RBAC_ROLES]: The mapped internal RBAC role if a mapping exists, 

303 None otherwise 

304 

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 

311 

312 jwt_role: Final = self.get_jwt_role(token=token, default_value=None) 

313 if not jwt_role: 

314 return None 

315 

316 jwt_role_set: Final = set(jwt_role) 

317 

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 

322 

323 return None 

324 

325 def get_rbac_role(self, token: dict) -> RBAC_ROLES | None: 

326 """ 

327 Returns the RBAC role the token 'belongs' to. 

328 

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 

333 

334 Resolves: https://github.com/BerriAI/litellm/issues/6793 

335 

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) 

345 

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 

358 

359 return None 

360 

361 def is_admin(self, scopes: list) -> bool: 

362 if self.litellm_jwtauth.admin_jwt_scope in scopes: 

363 return True 

364 return False 

365 

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 

370 

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) 

374 

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 

377 

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 [] 

390 

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 [] 

398 

399 return [] 

400 

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. 

406 

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. 

412 

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 

439 

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) 

443 

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 

455 

456 return user_id 

457 

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 

467 

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

474 

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 

480 

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 

487 

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 

524 

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. 

528 

529 Args: 

530 token: The decoded JWT token dictionary 

531 default_value: Default value to return if field not found 

532 

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 

549 

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 

559 

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) 

563 

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 

576 

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. 

580 

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 

595 

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 

600 

601 jwt_roles: Final = self.get_jwt_role(token=token, default_value=[]) 

602 if not jwt_roles: 

603 return None 

604 

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 

610 

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. 

614 

615 Returns the jwt role from the token. 

616 

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 

631 

632 def is_allowed_user_role(self, user_roles: list[str] | None) -> bool: 

633 """ 

634 Returns the user role from the token. 

635 

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 

645 

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) 

649 

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 

662 

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 

676 

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 

682 

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) 

686 

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 

699 

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. 

703 

704 Args: 

705 token: The decoded JWT token dictionary 

706 default_value: Default value to return if field not found 

707 

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 

724 

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 

737 

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 

747 

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 ) 

754 

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) 

770 

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 

775 

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 

779 

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 

784 

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) 

787 

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 

799 

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 ) 

808 

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. 

817 

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 

834 

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 

850 

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 

854 

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. 

863 

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) 

871 

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 

879 

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 

889 

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

901 

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

905 

906 verbose_proxy_logger.debug("JWT Auth: Resolved OIDC discovery %s -> jwks_uri=%s", url, jwks_uri) 

907 return jwks_uri 

908 

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 

914 

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 

920 

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 ) 

927 

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

933 

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 

936 

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 ) 

945 

946 public_key: Final = self.parse_keys(keys=keys, kid=kid) 

947 if public_key is not None: 

948 return cast(dict, public_key) 

949 

950 raise NoMatchingJWTPublicKeyError(f"No matching public key found. keys={resolved_jwks_url}, kid={kid}") 

951 

952 async def get_public_key(self, kid: str | None) -> dict: 

953 keys_url: Final = os.getenv("JWT_PUBLIC_KEY_URL") 

954 

955 if keys_url is None: 

956 raise Exception("Missing JWT Public Key URL from environment.") 

957 

958 keys_url_list: Final = [url.strip() for url in keys_url.split(",") if url.strip()] 

959 

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 

968 

969 raise NoMatchingJWTPublicKeyError(f"No matching public key found. keys={keys_url_list}, kid={kid}") 

970 

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 

986 

987 return public_key 

988 

989 def is_allowed_domain(self, user_email: str) -> bool: 

990 if self.litellm_jwtauth.user_allowed_email_domain is None: 

991 return True 

992 

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 

998 

999 async def get_oidc_userinfo(self, token: str) -> dict: 

1000 """ 

1001 Fetch user information from OIDC UserInfo endpoint. 

1002 

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. 

1006 

1007 Args: 

1008 token: The access token to use for authentication 

1009 

1010 Returns: 

1011 dict: User information from the UserInfo endpoint 

1012 

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

1018 

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) 

1022 

1023 if cached_userinfo is not None: 

1024 verbose_proxy_logger.debug("Returning cached OIDC UserInfo") 

1025 return cached_userinfo 

1026 

1027 verbose_proxy_logger.debug("Calling OIDC UserInfo endpoint: %s", self.litellm_jwtauth.oidc_userinfo_endpoint) 

1028 

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 ) 

1038 

1039 if response.status_code != 200: 

1040 raise Exception(f"OIDC UserInfo endpoint returned status {response.status_code}: {response.text}") 

1041 

1042 userinfo: Final = _userinfo_document(response) 

1043 verbose_proxy_logger.debug("Received OIDC UserInfo: %s", userinfo) 

1044 

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 ) 

1051 

1052 return userinfo 

1053 

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

1057 

1058 _unscoped_jwt_warning_emitted = False 

1059 

1060 @classmethod 

1061 def _build_decode_kwargs(cls) -> dict: 

1062 """Build the audience/issuer/options kwargs for ``jwt.decode``. 

1063 

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. 

1069 

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

1077 

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 

1086 

1087 options: Final[dict] = {} 

1088 if audience is None: 

1089 options["verify_aud"] = False 

1090 if issuer is None: 

1091 options["verify_iss"] = False 

1092 

1093 return { 

1094 "audience": audience, 

1095 "issuer": issuer, 

1096 "options": options or None, 

1097 } 

1098 

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 

1103 

1104 issuer_configs: Final = litellm_jwtauth.issuers 

1105 if not issuer_configs: 

1106 return None 

1107 

1108 claims: Final = self.get_unverified_claims(token=token) 

1109 if claims is None: 

1110 return None 

1111 

1112 issuer: Final = claims.get("iss") 

1113 if not isinstance(issuer, str) or not issuer: 

1114 return None 

1115 

1116 for issuer_config in issuer_configs: 

1117 if issuer_config.issuer == issuer: 

1118 return issuer_config 

1119 

1120 return None 

1121 

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" 

1128 

1129 def _get_claim_value_for_issuer_mapping(self, token: dict, claim_field: str) -> Any: 

1130 """Resolve a mapped claim from ``token``. 

1131 

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 

1145 

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 ] 

1157 

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 

1167 

1168 return normalized 

1169 

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 

1176 

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 

1197 

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 ) 

1216 

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 ) 

1228 

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 ) 

1243 

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 

1252 

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

1270 

1271 return self._apply_issuer_claim_mappings( 

1272 token=payload, 

1273 issuer_config=issuer_config, 

1274 ) 

1275 

1276 async def auth_jwt(self, token: str) -> dict: 

1277 header: Final = jwt.get_unverified_header(token) 

1278 

1279 verbose_proxy_logger.debug("header: %s", header) 

1280 

1281 kid: Final = header.get("kid", None) 

1282 

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 ) 

1290 

1291 decode_kwargs: Final = self._build_decode_kwargs() 

1292 

1293 public_key: Final = await self.get_public_key(kid=kid) 

1294 

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} 

1305 

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

1315 

1316 raise Exception("Invalid JWT Submitted") 

1317 

1318 async def close(self): 

1319 await self.http_handler.close() 

1320 

1321 

1322class JWTAuthManager: 

1323 """Manages JWT authentication and authorization operations""" 

1324 

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) 

1335 

1336 if role_based_routes is None or route is None: 

1337 return True 

1338 

1339 is_allowed: Final = _allowed_routes_check( 

1340 user_route=route, 

1341 allowed_routes=role_based_routes, 

1342 ) 

1343 

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 ) 

1349 

1350 return True 

1351 

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 

1364 

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 ) 

1374 

1375 return True 

1376 

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 

1389 

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) 

1394 

1395 requested_model: Final = request_data.get("model") 

1396 

1397 if not requested_model: 

1398 return 

1399 

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 

1408 

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 ) 

1435 

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 

1451 

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

1461 

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 ) 

1478 

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 

1497 

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) 

1511 

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 

1523 

1524 team_object: LiteLLM_TeamTable | None = None 

1525 

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 

1549 

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 

1565 

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 ) 

1595 

1596 return individual_team_id, team_object 

1597 

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) 

1602 

1603 all_team_ids: Final = set(team_ids_from_groups) 

1604 

1605 return all_team_ids 

1606 

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 

1620 

1621 if RouteChecks.jwt_team_routes_grant_pass_through(route=route, team_allowed_routes=team_allowed_routes): 

1622 return True 

1623 

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 ) 

1630 

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 ) 

1640 

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 

1655 

1656 denied_auth_enforced_pass_through_route = False 

1657 

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 

1668 

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 ) 

1679 

1680 if team_object is not None: 

1681 any_claim_team_resolved = True 

1682 

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 

1714 

1715 if denied_auth_enforced_pass_through_route: 

1716 JWTAuthManager._raise_team_passthrough_route_denial(route=route) 

1717 

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 ) 

1724 

1725 # No claim team resolved and fallback enabled — defer to fallback. 

1726 return None, None 

1727 

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 

1740 

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. 

1747 

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) 

1753 

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. 

1778 

1779 Returns ``(..., effective_user_id)``: JWT claim unless fuzzy lookup 

1780 matched a legacy row (GH #26789). 

1781 """ 

1782 

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 ) 

1810 

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 ) 

1819 

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 ) 

1841 

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 ) 

1856 

1857 return ( 

1858 user_object, 

1859 org_object, 

1860 end_user_object, 

1861 team_membership_object, 

1862 effective_user_id, 

1863 ) 

1864 

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 

1879 

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

1886 

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 ) 

1896 

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 

1920 

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. 

1939 

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 

1950 

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) 

1954 

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) 

1975 

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) 

1983 

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 

1995 

1996 if not user_object: 

1997 return 

1998 

1999 if not team_object: 

2000 return 

2001 

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 

2006 

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 

2034 

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 

2045 

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 

2049 

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 

2054 

2055 if user_object is None or prisma_client is None: 

2056 return 

2057 

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 ) 

2073 

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 ) 

2092 

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 

2107 

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. 

2121 

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 

2152 

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. 

2167 

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 

2177 

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 

2197 

2198 if not user_id: 

2199 return _tid, team_row, None 

2200 

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 

2210 

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. 

2232 

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). 

2238 

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. 

2242 

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 

2247 

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 

2296 

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 

2312 

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 ) 

2332 

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 ) 

2348 

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) 

2365 

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 ) 

2397 

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 

2409 

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 ) 

2423 

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) 

2462 

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 

2484 

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) 

2488 

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 ) 

2498 

2499 object_id = handler.get_object_id(token=jwt_valid_token, default_value=None) 

2500 

2501 # Get basic user info 

2502 user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token) 

2503 

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) 

2510 

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 

2516 

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 ) 

2522 

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} 

2568 

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) 

2573 

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) 

2589 

2590 header_db_fallback: Final = handler.litellm_jwtauth.fallback_to_db_teams and team_id is None 

2591 

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 ) 

2641 

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 ) 

2655 

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 ) 

2675 

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) 

2683 

2684 # Extract alias fields for resolution (if configured) 

2685 org_alias: Final = handler.get_org_alias(token=jwt_valid_token, default_value=None) 

2686 

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 ) 

2710 

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 

2713 

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 ) 

2722 

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 ) 

2785 

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 ) 

2792 

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 ) 

2800 

2801 # check if user is proxy admin 

2802 is_proxy_admin: Final = bool(user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN) 

2803 

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 ) 

2820 

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 )