Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/ui_sso.py: 11%
1816 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Has all /sso/* routes
4/sso/key/generate - handles user signing in with SSO and redirects to /sso/callback
5/sso/callback - returns JWT Redirect Response that redirects to LiteLLM UI
7/sso/debug/login - handles user signing in with SSO and redirects to /sso/debug/callback
8/sso/debug/callback - returns the OpenID object returned by the SSO provider
9"""
11import asyncio
12import base64
13import hashlib
14import inspect
15import json
16import os
17import re
18import secrets
19from collections.abc import Awaitable, Callable, Mapping, Sequence
20from copy import deepcopy
21from html import escape
22from types import MappingProxyType
23from typing import (
24 TYPE_CHECKING,
25 Any,
26 Final,
27 Literal,
28 NoReturn,
29 Optional,
30 Protocol,
31 TypeAlias,
32 Union,
33 cast,
34 overload,
35)
36from urllib.parse import parse_qs, urlencode, urlparse
38if TYPE_CHECKING: 38 ↛ 39line 38 didn't jump to line 39 because the condition on line 38 was never true
39 import httpx
41import jwt
42from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
43from fastapi.responses import RedirectResponse
44from pydantic import BaseModel, TypeAdapter, ValidationError
46import litellm
47from litellm._logging import verbose_proxy_logger
48from litellm._uuid import uuid
49from litellm.caching.dual_cache import DualCache
50from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
51from litellm.constants import (
52 CLI_SSO_CLAIM_MAP,
53 CLI_SSO_CLAIM_MAX_SCALAR_LENGTH,
54 CLI_SSO_SESSION_CACHE_KEY_PREFIX,
55 CLI_SSO_SESSION_TTL_SECONDS,
56 LITELLM_CLI_SOURCE_IDENTIFIER,
57 LITELLM_UI_SESSION_DURATION,
58 MAX_SPENDLOG_ROWS_TO_QUERY,
59 MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE,
60 MICROSOFT_USER_EMAIL_ATTRIBUTE,
61 MICROSOFT_USER_FIRST_NAME_ATTRIBUTE,
62 MICROSOFT_USER_ID_ATTRIBUTE,
63 MICROSOFT_USER_LAST_NAME_ATTRIBUTE,
64)
65from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
66from litellm.llms.custom_httpx.http_handler import (
67 AsyncHTTPHandler,
68 get_async_httpx_client,
69 httpxSpecialProvider,
70)
71from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
72 SSOIdentityAssertion,
73 assertion_from_sso_login,
74 ema_assertion_retention_enabled,
75 retain_sso_identity_assertion_for_ema,
76)
77from litellm.proxy._types import (
78 CommonProxyErrors,
79 LiteLLM_UserTable,
80 LitellmUserRoles,
81 Member,
82 NewTeamRequest,
83 NewUserRequest,
84 NewUserResponse,
85 ProxyErrorTypes,
86 ProxyException,
87 SSOUserDefinedValues,
88 TeamMemberAddRequest,
89 UserAPIKeyAuth,
90)
91from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object
92from litellm.proxy.auth.auth_utils import (
93 _get_request_ip_address,
94 has_user_setup_sso,
95)
96from litellm.proxy.auth.handle_jwt import JWTHandler
97from litellm.proxy.auth.ip_address_utils import IPAddressUtils
98from litellm.proxy.auth.team_grants import TeamModelAliasTable
99from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
100from litellm.proxy.common_utils.admin_ui_utils import (
101 admin_ui_disabled,
102 show_missing_vars_in_env,
103)
104from litellm.proxy.common_utils.html_forms.default_credentials_hint import should_hide_default_credentials_hint
105from litellm.proxy.common_utils.html_forms.jwt_display_template import (
106 jwt_display_template,
107)
108from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
109from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
110from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
111from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
112from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
113 id_jag_assertion_capture_gap,
114)
115from litellm.proxy.management_endpoints.sso.saml_sso import SAMLAuthHandler
116from litellm.proxy.management_endpoints.sso_helper_utils import (
117 check_is_admin_only_access,
118 has_admin_ui_access,
119)
120from litellm.proxy.management_endpoints.team_endpoints import new_team, team_member_add
121from litellm.proxy.management_endpoints.types import (
122 LITELLM_USER_ROLE_HIERARCHY,
123 CustomOpenID,
124 get_litellm_user_role,
125 is_valid_litellm_user_role,
126)
127from litellm.proxy.utils import (
128 PrismaClient,
129 ProxyLogging,
130 get_custom_url,
131 get_server_root_path,
132)
133from litellm.repositories.prisma_protocols import TableActions
134from litellm.repositories.table_repositories import SSOConfigRepository
135from litellm.repositories.team_repository import TeamRepository
136from litellm.repositories.user_repository import UserRepository
137from litellm.secret_managers.main import get_secret_bool, get_secret_str, str_to_bool
138from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403
139from litellm.types.proxy.management_endpoints.ui_sso import (
140 DefaultTeamSSOParams,
141 MicrosoftGraphAPIUserGroupDirectoryObject,
142 MicrosoftGraphAPIUserGroupResponse,
143 MicrosoftServicePrincipalTeam,
144 RoleMappings,
145 TeamMappings,
146)
147from litellm.types.proxy.ui_sso import ParsedOpenIDResult
149if TYPE_CHECKING: 149 ↛ 150line 149 didn't jump to line 150 because the condition on line 149 was never true
150 from fastapi_sso.sso.base import OpenID
151else:
152 from typing import Any as OpenID
154router: Final = APIRouter()
156# OAuth bearer credential fields that must not appear in SSO debug responses
157# (received_response is included in restricted-group error messages).
158# Metadata fields (token_type, expires_in, scope) are intentionally kept so
159# response convertors see the same fields in the PKCE path as in the non-PKCE path.
160_OAUTH_TOKEN_FIELDS: Final = frozenset({"access_token", "id_token", "refresh_token"})
161_CLI_SSO_FLOW_CACHE_KEY_PREFIX: Final = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:flow"
162_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX: Final = f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:start_rate_limit"
163_CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS: Final = 60
164_CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS: Final = 30
165_CLI_SSO_USER_CODE_ALPHABET: Final = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
166_CLI_SSO_LOGIN_ID_RE: Final = re.compile(r"^cli-[A-Za-z0-9_-]{12,124}$")
167_CLI_SSO_USER_CODE_RE = re.compile(rf"^[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}-[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}$")
168_CLI_SSO_SCALAR_TYPES: Final = (str, int, float, bool)
169_CLI_SSO_DEST_KEY_RE: Final = re.compile(r"^[A-Za-z0-9_.-]+$")
170_CLI_SSO_SECRET_KEY_FRAGMENTS: Final = frozenset(
171 {
172 "access_token",
173 "api_key",
174 "client_secret",
175 "id_token",
176 "password",
177 "private_key",
178 "refresh_token",
179 "secret",
180 }
181)
184class _UserMetadataRow(Protocol):
185 @property
186 def metadata(self) -> Mapping[str, object] | None: ... 186 ↛ exitline 186 didn't return from function 'metadata' because
189def _user_meta_db(repo: UserRepository) -> "TableActions[_UserMetadataRow]":
190 return cast( # cast-ok: prisma types Json columns as str; the client hands back the deserialized value
191 "TableActions[_UserMetadataRow]", repo.table
192 )
195class _SsoConfigRow(Protocol):
196 @property
197 def sso_settings(self) -> Mapping[str, object] | None: ... 197 ↛ exitline 197 didn't return from function 'sso_settings' because
200def _sso_config_db(repo: SSOConfigRepository) -> "TableActions[_SsoConfigRow]":
201 return cast( # cast-ok: prisma types Json columns as str; the client hands back the deserialized value
202 "TableActions[_SsoConfigRow]", repo.table
203 )
206class _TeamDetailRow(Protocol):
207 def model_dump(self) -> Mapping[str, object]: ... 207 ↛ exitline 207 didn't return from function 'model_dump' because
210def _team_detail_db(repo: TeamRepository) -> "TableActions[_TeamDetailRow]":
211 return repo.table
214_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
217class _TeamRowGrants(BaseModel):
218 team_id: str
219 team_alias: str | None = None
220 models: tuple[str, ...] = ()
221 litellm_model_table: TeamModelAliasTable | None = None
224class CliSsoTeamDetail(BaseModel):
225 """The per-team snapshot cached in the CLI SSO flow and echoed to the CLI on poll."""
227 team_id: str | None = None
228 team_alias: str | None = None
229 team_models: tuple[str, ...]
230 team_model_aliases: Mapping[str, str] | None = None
233_CLI_SSO_TEAM_DETAILS_ADAPTER: Final = TypeAdapter(tuple[CliSsoTeamDetail, ...])
234_TEAMLESS_CLI_SSO_TEAM_DETAIL: Final = CliSsoTeamDetail(team_models=())
237class _CustomSsoCall(Protocol):
238 async def __call__(self, sso_response: object) -> SSOUserDefinedValues | None: ... 238 ↛ exitline 238 didn't return from function '__call__' because
241class _ServicePrincipalAssignment(Protocol):
242 def get(self, key: str) -> str: ... 242 ↛ exitline 242 didn't return from function 'get' because
245class _ServicePrincipalPage(Protocol):
246 @overload
247 def get( 247 ↛ exitline 247 didn't return from function 'get' because
248 self,
249 key: Literal["value"],
250 default: Sequence["_ServicePrincipalAssignment"],
251 ) -> Sequence["_ServicePrincipalAssignment"]: ...
253 @overload
254 def get(self, key: Literal["@odata.nextLink"]) -> str | None: ... 254 ↛ exitline 254 didn't return from function 'get' because
257def _as_object(value: object) -> object:
258 return value
261def _hash_cli_sso_secret(secret: str) -> str:
262 return hashlib.sha256(secret.encode("utf-8")).hexdigest()
265def _normalize_cli_sso_user_code(user_code: str) -> str:
266 return "".join(ch for ch in user_code.upper() if ch.isalnum())
269def _generate_cli_sso_user_code() -> str:
270 user_code: Final = "".join(secrets.choice(_CLI_SSO_USER_CODE_ALPHABET) for _ in range(8))
271 return f"{user_code[:4]}-{user_code[4:]}"
274def _get_cli_sso_flow_cache_key(login_id: str) -> str:
275 return f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:{login_id}"
278def _is_valid_cli_sso_login_id(login_id: str | None) -> bool:
279 return isinstance(login_id, str) and bool(_CLI_SSO_LOGIN_ID_RE.fullmatch(login_id))
282def _is_valid_cli_sso_user_code(user_code: str | None) -> bool:
283 return isinstance(user_code, str) and bool(_CLI_SSO_USER_CODE_RE.fullmatch(user_code))
286def _cli_sso_verification_uri_complete_enabled() -> bool:
287 from litellm.proxy.proxy_server import general_settings
289 return bool(general_settings.get("allow_cli_sso_verification_uri_complete", False))
292def _cli_sso_start_response_body(
293 *,
294 login_id: str,
295 poll_secret: str,
296 user_code: str,
297 verification_uri_complete: str | None,
298) -> dict[str, str | int]:
299 if verification_uri_complete is None:
300 return {
301 "login_id": login_id,
302 "poll_secret": poll_secret,
303 "user_code": user_code,
304 "expires_in": CLI_SSO_SESSION_TTL_SECONDS,
305 }
306 return {
307 "login_id": login_id,
308 "poll_secret": poll_secret,
309 "user_code": user_code,
310 "verification_uri_complete": verification_uri_complete,
311 "expires_in": CLI_SSO_SESSION_TTL_SECONDS,
312 }
315def _get_cli_sso_start_rate_limit_cache_key(request: Request, use_x_forwarded_for: bool | None = False) -> str:
316 client_ip: Final = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) or "unknown"
317 client_ip_hash: Final = _hash_cli_sso_secret(client_ip)
318 return f"{_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX}:{client_ip_hash}"
321def _check_cli_sso_start_rate_limit(
322 request: Request,
323 cache: DualCache,
324 use_x_forwarded_for: bool | None = False,
325) -> None:
326 rate_limit_cache_key: Final = _get_cli_sso_start_rate_limit_cache_key(
327 request=request, use_x_forwarded_for=use_x_forwarded_for
328 )
329 current_attempts: Final = cache.increment_cache(
330 key=rate_limit_cache_key,
331 value=1,
332 ttl=_CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS,
333 )
334 if current_attempts > _CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS:
335 raise HTTPException(
336 status_code=429,
337 detail="Too many CLI login attempts. Try again later.",
338 )
341def _read_cli_sso_flow(cache: DualCache, cache_key: str) -> object:
342 redis_cache: Final = cache.redis_cache
343 if redis_cache is None:
344 return cache.get_cache(key=cache_key)
345 try:
346 return redis_cache.get_cache(key=cache_key)
347 except RedisCircuitBreakerOpenError:
348 return None
351def _get_cli_sso_flow_or_raise(login_id: str | None, cache: DualCache) -> dict:
352 if isinstance(login_id, str) and login_id.startswith("sk-"):
353 raise HTTPException(
354 status_code=400,
355 detail=(
356 "Your litellm CLI is out of date and uses a login flow this proxy no longer supports. "
357 "Upgrade it with `pip install -U 'litellm[proxy]'` and run `lite login` again."
358 ),
359 )
360 if not _is_valid_cli_sso_login_id(login_id):
361 raise HTTPException(status_code=400, detail="Invalid CLI login session id")
363 flow = _read_cli_sso_flow(cache, _get_cli_sso_flow_cache_key(cast(str, login_id)))
364 if isinstance(flow, str):
365 try:
366 flow = _as_object(json.loads(flow))
367 except ValueError:
368 flow = None
369 if not isinstance(flow, dict) or "poll_secret_hash" not in flow:
370 verbose_proxy_logger.warning(
371 "CLI SSO login session not found in cache for login_id=%s. If the proxy runs multiple replicas, "
372 "a shared Redis cache is required for CLI login to work.",
373 login_id,
374 )
375 raise HTTPException(
376 status_code=400,
377 detail=(
378 "CLI login session not found or expired. Run `lite login` again. "
379 "If this happens immediately after starting a login, the proxy is likely running multiple "
380 "replicas without a shared cache; configure a Redis cache "
381 "so every replica can see the login session."
382 ),
383 )
384 return flow
387def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None:
388 cache_key: Final = _get_cli_sso_flow_cache_key(login_id)
389 redis_cache: Final = cache.redis_cache
390 if redis_cache is not None:
391 redis_cache.set_cache(key=cache_key, value=json.dumps(flow), ttl=CLI_SSO_SESSION_TTL_SECONDS)
392 else:
393 cache.set_cache(key=cache_key, value=flow, ttl=CLI_SSO_SESSION_TTL_SECONDS)
396def _verify_cli_sso_poll_secret(flow: dict, poll_secret: str | None) -> bool:
397 expected_poll_secret_hash: Final = flow.get("poll_secret_hash")
398 if not isinstance(expected_poll_secret_hash, str) or not isinstance(poll_secret, str):
399 return False
400 supplied_poll_secret_hash: Final = _hash_cli_sso_secret(poll_secret)
401 return secrets.compare_digest(supplied_poll_secret_hash, expected_poll_secret_hash)
404def _parse_cli_sso_claim_map() -> list[tuple[str, str]]:
405 """
406 Parse CLI_SSO_CLAIM_MAP / LITELLM_CLI_SSO_CLAIM_MAP.
408 Format: comma-separated ``source_claim->metadata_key`` pairs, e.g.
409 ``employment_type->acme_employment_type,org_info.department->department``.
410 Destination keys may use an optional ``metadata.`` prefix; values are stored
411 on the LiteLLM user's ``metadata`` JSON column.
412 """
413 claim_map_raw: Final = CLI_SSO_CLAIM_MAP.strip()
414 if not claim_map_raw:
415 return []
417 parsed: Final[list[tuple[str, str]]] = []
418 for entry in claim_map_raw.split(","):
419 entry = entry.strip()
420 if not entry or "->" not in entry:
421 continue
422 source_claim, dest_key = entry.split("->", 1)
423 source_claim = source_claim.strip()
424 dest_key = dest_key.strip()
425 dest_key = dest_key.removeprefix("metadata.")
426 if source_claim and dest_key:
427 parsed.append((source_claim, dest_key))
428 return parsed
431def _is_safe_cli_sso_metadata_dest_key(dest_key: str) -> bool:
432 if not dest_key or not _CLI_SSO_DEST_KEY_RE.fullmatch(dest_key):
433 return False
434 lowered: Final = dest_key.lower()
435 return not any(fragment in lowered for fragment in _CLI_SSO_SECRET_KEY_FRAGMENTS)
438def _is_safe_cli_sso_scalar_claim_value(value: object) -> bool:
439 if not isinstance(value, _CLI_SSO_SCALAR_TYPES):
440 return False
441 if isinstance(value, str):
442 if len(value) > CLI_SSO_CLAIM_MAX_SCALAR_LENGTH:
443 return False
444 if value.startswith("eyJ") and value.count(".") >= 2:
445 return False
446 return True
449def _sso_result_to_dict(result: CustomOpenID | OpenID | dict[str, object]) -> dict[str, object]:
450 if isinstance(result, dict):
451 return result
452 if hasattr(result, "model_dump"):
453 dumped: Final = result.model_dump()
454 if isinstance(dumped, dict):
455 return dumped
456 return {}
459def _get_nested_claim_value(data: Mapping[str, object], claim_path: str) -> object:
460 """Resolve a dot-notation claim path against an SSO result dict.
462 Unlike ``get_nested_value``, this does not strip a leading ``metadata.``
463 prefix, since OIDC claims may legitimately use ``metadata`` as a top-level
464 key.
465 """
466 if not claim_path:
467 return None
468 if claim_path in data:
469 return data[claim_path]
470 placeholder: Final = "\x00"
471 parts = claim_path.replace("\\.", placeholder).split(".")
472 parts = [p.replace(placeholder, ".") for p in parts]
473 current: object = data
474 for part in parts:
475 if isinstance(current, dict) and part in current:
476 current = current[part]
477 else:
478 return None
479 return current
482def _extract_sso_claim_value(result: CustomOpenID | OpenID | dict[str, object], claim_path: str) -> object:
483 extra_fields: Final = getattr(result, "extra_fields", None)
484 if isinstance(extra_fields, dict):
485 if claim_path in extra_fields:
486 return extra_fields[claim_path]
487 nested: Final = _get_nested_claim_value(extra_fields, claim_path)
488 if nested is not None:
489 return nested
491 if isinstance(result, dict):
492 return _get_nested_claim_value(result, claim_path)
494 result_dict: Final = _sso_result_to_dict(result)
495 return _get_nested_claim_value(result_dict, claim_path)
498def _set_nested_metadata_value(metadata: dict[str, object], key_path: str, value: object) -> None:
499 placeholder: Final = "\x00"
500 parts = key_path.replace("\\.", placeholder).split(".")
501 parts = [p.replace(placeholder, ".") for p in parts]
502 current: dict[str, object] = metadata
503 for part in parts[:-1]:
504 existing = current.get(part)
505 if not isinstance(existing, dict):
506 existing = {}
507 current[part] = existing
508 current = existing
509 current[parts[-1]] = value
512def _flatten_cli_sso_metadata_for_poll(
513 metadata: Mapping[str, object],
514) -> dict[str, str | int | float | bool]:
515 """Expose scalar attribution metadata as a flat dict for CLI poll responses."""
516 flattened: Final[dict[str, str | int | float | bool]] = {}
517 stack: Final[list[tuple[str, object]]] = [("", metadata)]
518 while stack:
519 prefix, value = stack.pop()
520 if isinstance(value, dict):
521 nested_items: Mapping[str, object] = value
522 for key, nested in nested_items.items():
523 nested_prefix = f"{prefix}.{key}" if prefix else key
524 stack.append((nested_prefix, nested))
525 elif isinstance(value, (str, int, float, bool)) and _is_safe_cli_sso_scalar_claim_value(value):
526 flattened[prefix] = value
527 return flattened
530def build_cli_sso_attribution_metadata(
531 result: CustomOpenID | OpenID | dict[str, object],
532) -> dict[str, object]:
533 """
534 Build allowlisted, non-secret scalar attribution metadata from an SSO result.
536 Sources are configured via CLI_SSO_CLAIM_MAP / LITELLM_CLI_SSO_CLAIM_MAP and
537 may include claims captured by GENERIC_USER_EXTRA_ATTRIBUTES on CustomOpenID.
538 """
539 claim_map: Final = _parse_cli_sso_claim_map()
540 if not claim_map:
541 return {}
543 metadata: Final[dict[str, object]] = {}
544 for source_claim, dest_key in claim_map:
545 if not _is_safe_cli_sso_metadata_dest_key(dest_key):
546 verbose_proxy_logger.debug("Skipping unsafe CLI SSO metadata destination key: %s", dest_key)
547 continue
549 raw_value = _extract_sso_claim_value(result=result, claim_path=source_claim)
550 if not _is_safe_cli_sso_scalar_claim_value(raw_value):
551 continue
553 _set_nested_metadata_value(metadata=metadata, key_path=dest_key, value=raw_value)
555 return metadata
558def _merge_cli_sso_attribution_metadata(
559 existing_metadata: dict[str, object], attribution_metadata: dict[str, object]
560) -> dict[str, object]:
561 """Merge attribution metadata into existing user metadata in-place.
563 Preserves original value types (in particular, string claim values that
564 happen to look numeric are NOT coerced to ``int``/``float``). Nested dicts
565 are merged iteratively so attribution claims do not clobber unrelated keys
566 under the same parent.
567 """
568 pending: Final[list[tuple[dict[str, object], dict[str, object]]]] = [(existing_metadata, attribution_metadata)]
569 while pending:
570 target, source = pending.pop()
571 for key, value in source.items():
572 if value is None:
573 continue
574 existing_value = target.get(key)
575 if isinstance(value, dict) and isinstance(existing_value, dict):
576 pending.append((existing_value, value))
577 else:
578 target[key] = value
579 return existing_metadata
582async def _persist_cli_sso_user_metadata(
583 prisma_client: PrismaClient,
584 user_id: str,
585 attribution_metadata: dict[str, object],
586) -> None:
587 if not attribution_metadata:
588 return
590 try:
591 user_row: Final = await _user_meta_db(UserRepository(prisma_client)).find_unique(where={"user_id": user_id})
592 existing_metadata: dict[str, object] = {}
593 if user_row is not None:
594 row_metadata: Final = user_row.metadata
595 if isinstance(row_metadata, dict):
596 existing_metadata = deepcopy(row_metadata)
598 merged_metadata: Final = _merge_cli_sso_attribution_metadata(
599 existing_metadata=existing_metadata,
600 attribution_metadata=attribution_metadata,
601 )
602 await _user_meta_db(UserRepository(prisma_client)).update_many(
603 where={"user_id": user_id},
604 data={"metadata": merged_metadata},
605 )
606 verbose_proxy_logger.info(
607 "Persisted CLI SSO attribution metadata for user %s: %s",
608 user_id,
609 list(_flatten_cli_sso_metadata_for_poll(attribution_metadata).keys()),
610 )
611 except Exception as e:
612 verbose_proxy_logger.error("Failed to persist CLI SSO attribution metadata for user %s: %s", user_id, e)
615def _cli_poll_attribution_metadata_from_session(
616 session_data: Mapping[str, object],
617) -> dict[str, str | int | float | bool]:
618 stored: Final = session_data.get("attribution_metadata")
619 if isinstance(stored, dict):
620 return _flatten_cli_sso_metadata_for_poll(stored)
621 return {}
624def _render_cli_sso_verification_page(
625 verify_url: str,
626 browser_complete_token: str,
627 prefill_user_code: str | None = None,
628) -> str:
629 escaped_verify_url: Final = escape(verify_url, quote=True)
630 escaped_browser_complete_token: Final = escape(browser_complete_token, quote=True)
631 user_code_value_attr: Final = f' value="{escape(prefill_user_code, quote=True)}"' if prefill_user_code else ""
632 instructions: Final = (
633 "Confirm the verification code below to finish this login."
634 if prefill_user_code
635 else "Enter the verification code shown in your terminal to finish this login."
636 )
637 return f"""
638 <!doctype html>
639 <html>
640 <head>
641 <title>LiteLLM CLI Login</title>
642 <style>
643 body {{
644 font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif;
645 margin: 0;
646 min-height: 100vh;
647 display: flex;
648 align-items: center;
649 justify-content: center;
650 background: #f8fafc;
651 color: #0f172a;
652 }}
653 main {{
654 width: min(420px, calc(100vw - 32px));
655 background: #ffffff;
656 border: 1px solid #e2e8f0;
657 border-radius: 8px;
658 padding: 28px;
659 box-shadow: 0 12px 32px rgba(15, 23, 42, 0.08);
660 }}
661 h1 {{ font-size: 22px; margin: 0 0 12px; }}
662 p {{ line-height: 1.5; margin: 0 0 18px; color: #334155; }}
663 label {{ display: block; font-weight: 600; margin-bottom: 8px; }}
664 input {{
665 box-sizing: border-box;
666 width: 100%;
667 padding: 12px;
668 border: 1px solid #cbd5e1;
669 border-radius: 6px;
670 font-size: 20px;
671 letter-spacing: 0.08em;
672 text-transform: uppercase;
673 }}
674 button {{
675 margin-top: 16px;
676 width: 100%;
677 padding: 12px;
678 border: 0;
679 border-radius: 6px;
680 background: #0f172a;
681 color: #ffffff;
682 font-weight: 600;
683 cursor: pointer;
684 }}
685 </style>
686 </head>
687 <body>
688 <main>
689 <h1>Complete CLI Login</h1>
690 <p>{instructions}</p>
691 <form method="post" action="{escaped_verify_url}">
692 <input type="hidden" name="browser_complete_token" value="{escaped_browser_complete_token}" />
693 <label for="user_code">Verification code</label>
694 <input id="user_code" name="user_code" autocomplete="one-time-code"{user_code_value_attr} required autofocus />
695 <button type="submit">Continue</button>
696 </form>
697 </main>
698 </body>
699 </html>
700 """
703@router.post("/sso/cli/start", tags=["experimental"], include_in_schema=False)
704async def cli_sso_start(request: Request):
705 from litellm.proxy.proxy_server import cli_sso_session_cache, general_settings
707 _check_cli_sso_start_rate_limit(
708 request=request,
709 cache=cli_sso_session_cache,
710 use_x_forwarded_for=bool((general_settings or {}).get("use_x_forwarded_for", False)),
711 )
713 login_id: Final = f"cli-{secrets.token_urlsafe(24)}"
714 poll_secret: Final = secrets.token_urlsafe(32)
715 user_code: Final = _generate_cli_sso_user_code()
717 flow: Final = {
718 "poll_secret_hash": _hash_cli_sso_secret(poll_secret),
719 "user_code_hash": _hash_cli_sso_secret(_normalize_cli_sso_user_code(user_code)),
720 "sso_complete": False,
721 "user_code_verified": False,
722 "session_data": None,
723 }
724 _set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow)
726 verification_uri_complete: Final[str | None] = (
727 (
728 get_custom_url(request_base_url=str(request.base_url), route="sso/key/generate")
729 + "?"
730 + urlencode(
731 {
732 "source": LITELLM_CLI_SOURCE_IDENTIFIER,
733 "key": login_id,
734 "user_code": user_code,
735 }
736 )
737 )
738 if _cli_sso_verification_uri_complete_enabled()
739 else None
740 )
741 return _cli_sso_start_response_body(
742 login_id=login_id,
743 poll_secret=poll_secret,
744 user_code=user_code,
745 verification_uri_complete=verification_uri_complete,
746 )
749@router.post("/sso/cli/complete/{login_id}", tags=["experimental"], include_in_schema=False)
750async def cli_sso_complete(request: Request, login_id: str):
751 from fastapi.responses import HTMLResponse
753 from litellm.proxy.common_utils.html_forms.cli_sso_success import (
754 render_cli_sso_success_page,
755 )
756 from litellm.proxy.proxy_server import cli_sso_session_cache
758 flow: Final = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cli_sso_session_cache)
759 if not flow.get("sso_complete") or not flow.get("session_data"):
760 raise HTTPException(status_code=400, detail="CLI login is not ready")
762 body: Final = (await request.body()).decode("utf-8")
763 form_values: Final = parse_qs(body)
764 supplied_user_code: Final = (form_values.get("user_code") or [""])[0]
765 supplied_browser_complete_token: Final = (form_values.get("browser_complete_token") or [""])[0]
766 supplied_user_code_hash: Final = _hash_cli_sso_secret(_normalize_cli_sso_user_code(supplied_user_code))
767 supplied_browser_complete_token_hash: Final = _hash_cli_sso_secret(supplied_browser_complete_token)
769 expected_user_code_hash: Final = flow.get("user_code_hash")
770 if not isinstance(expected_user_code_hash, str) or not secrets.compare_digest(
771 supplied_user_code_hash, expected_user_code_hash
772 ):
773 raise HTTPException(status_code=400, detail="Invalid verification code")
775 expected_browser_complete_token_hash: Final = flow.get("browser_complete_token_hash")
776 if not isinstance(expected_browser_complete_token_hash, str) or not secrets.compare_digest(
777 supplied_browser_complete_token_hash, expected_browser_complete_token_hash
778 ):
779 raise HTTPException(status_code=400, detail="Invalid verification code")
781 flow["user_code_verified"] = True
782 _set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow)
784 html_content: Final = render_cli_sso_success_page()
785 return HTMLResponse(content=html_content, status_code=200)
788def normalize_email(email: str | None) -> str | None:
789 """
790 Normalize email address to lowercase for consistent storage and comparison.
792 Email addresses should be treated as case-insensitive for SSO purposes,
793 even though RFC 5321 technically allows case-sensitive local parts.
794 This prevents issues where SSO providers return emails with different casing
795 than what's stored in the database.
797 Args:
798 email: Email address to normalize, can be None
800 Returns:
801 Lowercased email address, or None if input is None
802 """
803 if email is None:
804 return None
805 return email.lower() if isinstance(email, str) else email
808def determine_role_from_groups(
809 user_groups: list[str],
810 role_mappings: "RoleMappings",
811) -> LitellmUserRoles | None:
812 """
813 Determine the highest privilege role for a user based on their groups.
815 Role hierarchy (highest to lowest):
816 - proxy_admin
817 - proxy_admin_viewer
818 - internal_user
819 - internal_user_viewer
821 Args:
822 user_groups: List of group names from the SSO token
823 role_mappings: RoleMappings configuration object
825 Returns:
826 The highest privilege role found, or default_role if no matches, or None
827 """
828 if not role_mappings.roles:
829 # No role mappings configured, return default_role
830 return role_mappings.default_role
832 # Convert user_groups to a set for efficient lookup
833 user_groups_set: Final = set(user_groups) if isinstance(user_groups, list) else set()
835 # Find the highest privilege role the user belongs to
836 for role in LITELLM_USER_ROLE_HIERARCHY:
837 if role in role_mappings.roles:
838 role_groups = role_mappings.roles[role]
839 if isinstance(role_groups, list) and user_groups_set.intersection(set(role_groups)):
840 verbose_proxy_logger.debug(
841 "User groups %s matched role '%s' via groups: %s", user_groups, role.value, role_groups
842 )
843 return role
845 # No matching groups found, return default_role
846 verbose_proxy_logger.debug(
847 "User groups %s did not match any role mappings, using default_role: %s",
848 user_groups,
849 role_mappings.default_role,
850 )
851 return role_mappings.default_role
854def process_sso_jwt_access_token(
855 access_token_str: str | None,
856 sso_jwt_handler: JWTHandler | None,
857 result: OpenID | dict | None,
858 role_mappings: Optional["RoleMappings"] = None,
859) -> dict | None:
860 """
861 Process SSO JWT access token and extract team IDs and user role if available.
863 This function decodes the JWT access token and extracts team IDs and user
864 role, then sets them on the result object. Role extraction from the access
865 token is needed because some SSO providers (e.g., Keycloak) do not include
866 role claims in the UserInfo endpoint response.
868 Args:
869 access_token_str: The JWT access token string
870 sso_jwt_handler: SSO-specific JWT handler for team ID extraction
871 result: The SSO result object to update with team IDs and role
872 role_mappings: Optional role mappings configuration for group-based role determination
874 Returns:
875 The decoded access token payload dict, or None if decoding failed or
876 inputs were missing. Callers can pass this to _sync_user_role_from_jwt_role_map
877 so it has access to custom role claims (e.g. custom_roles) that are
878 encoded inside the JWT but stripped from received_response.
879 """
880 if access_token_str and result:
881 import jwt
883 try:
884 access_token_payload: Final = jwt.decode(access_token_str, options={"verify_signature": False})
885 except jwt.exceptions.DecodeError:
886 verbose_proxy_logger.debug(
887 "Access token is not a valid JWT (possibly an opaque token), skipping JWT-based extraction"
888 )
889 return None
891 # Extract team IDs from access token if sso_jwt_handler is available
892 if sso_jwt_handler:
893 if isinstance(result, dict):
894 result_team_ids: list[str] | None = result.get("team_ids", [])
895 if not result_team_ids:
896 team_ids = sso_jwt_handler.get_team_ids_from_jwt(access_token_payload)
897 result["team_ids"] = team_ids
898 else:
899 result_team_ids = getattr(result, "team_ids", []) if result else []
900 if not result_team_ids:
901 team_ids = sso_jwt_handler.get_team_ids_from_jwt(access_token_payload)
902 setattr(result, "team_ids", team_ids)
904 # Extract user role from access token if not already set from UserInfo
905 existing_role = result.get("user_role") if isinstance(result, dict) else getattr(result, "user_role", None)
906 if existing_role is None:
907 user_role: LitellmUserRoles | None = None
909 # Try role_mappings first (group-based role determination)
910 if role_mappings is not None and role_mappings.roles:
911 group_claim: Final = role_mappings.group_claim
912 user_groups_raw: Final[object] = get_nested_value(access_token_payload, group_claim)
914 user_groups: list[str] = []
915 if isinstance(user_groups_raw, list):
916 raw_groups: Final[Sequence[object]] = user_groups_raw
917 user_groups = [str(g) for g in raw_groups]
918 elif isinstance(user_groups_raw, str):
919 user_groups = [g.strip() for g in user_groups_raw.split(",") if g.strip()]
920 elif user_groups_raw is not None:
921 user_groups = [str(user_groups_raw)]
923 if user_groups:
924 user_role = determine_role_from_groups(user_groups, role_mappings)
925 verbose_proxy_logger.debug(
926 "Determined role '%s' from access token groups '%s' using role_mappings", user_role, user_groups
927 )
928 elif role_mappings.default_role:
929 user_role = role_mappings.default_role
931 # Fallback: try GENERIC_USER_ROLE_ATTRIBUTE on the access token payload
932 if user_role is None:
933 generic_user_role_attribute_name: Final = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role")
934 user_role_from_token: Final = get_nested_value(access_token_payload, generic_user_role_attribute_name)
935 if user_role_from_token is not None:
936 user_role = get_litellm_user_role(user_role_from_token)
937 verbose_proxy_logger.debug(
938 "Extracted role '%s' from access token field '%s'", user_role, generic_user_role_attribute_name
939 )
941 if user_role is not None:
942 if isinstance(result, dict):
943 result["user_role"] = user_role
944 else:
945 setattr(result, "user_role", user_role)
946 verbose_proxy_logger.debug("Set user_role='%s' from JWT access token", user_role)
948 return access_token_payload
950 return None
953def _decode_sso_token_claims(token: str | None) -> Mapping[str, object]:
954 if not token:
955 return MappingProxyType({})
956 try:
957 return MappingProxyType(
958 _SSO_TOKEN_CLAIMS_ADAPTER.validate_python(jwt.decode(token, options={"verify_signature": False}))
959 )
960 except (jwt.exceptions.InvalidTokenError, ValidationError):
961 verbose_proxy_logger.debug("SSO token is not a decodable JWT, skipping token claims")
962 return MappingProxyType({})
965def _merge_sso_token_claims(
966 userinfo: Mapping[str, object],
967 id_token: str | None,
968 access_token: str | None,
969) -> Mapping[str, object]:
970 sources: Final = (userinfo, _decode_sso_token_claims(id_token), _decode_sso_token_claims(access_token))
971 claim_names: Final = frozenset(key for source in sources for key in source)
972 return MappingProxyType(
973 {key: next((source[key] for source in sources if source.get(key) is not None), None) for key in claim_names}
974 )
977async def _raise_if_sso_exceeds_free_user_limit(premium_user: bool, prisma_client: PrismaClient | None) -> None:
978 """Free tier allows SSO for up to 5 billable users; beyond that requires an Enterprise license."""
979 if premium_user is True:
980 return
981 if prisma_client is None:
982 raise ProxyException(
983 message=CommonProxyErrors.db_not_connected_error.value,
984 type=ProxyErrorTypes.auth_error,
985 param="premium_user",
986 code=status.HTTP_403_FORBIDDEN,
987 )
988 billable_users: Final = await UserRepository(prisma_client).count_billable_users()
989 if billable_users and billable_users > 5:
990 raise ProxyException(
991 message="You must be a LiteLLM Enterprise user to use SSO for more than 5 users. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You configured SSO (one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, `GENERIC_CLIENT_ID`, or SAML) in your env. Please unset it",
992 type=ProxyErrorTypes.auth_error,
993 param="premium_user",
994 code=status.HTTP_403_FORBIDDEN,
995 )
998@router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False)
999async def google_login(
1000 request: Request,
1001 source: str | None = None,
1002 key: str | None = None,
1003 existing_key: str | None = None,
1004 return_to: str | None = None,
1005 user_code: str | None = None,
1006):
1007 """
1008 Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env
1009 PROXY_BASE_URL should be the your deployed proxy endpoint, e.g. PROXY_BASE_URL="https://litellm-production-7002.up.railway.app/"
1010 Example:
1011 """
1012 from litellm.proxy.proxy_server import (
1013 cli_sso_session_cache,
1014 general_settings,
1015 premium_user,
1016 prisma_client,
1017 user_api_key_cache,
1018 user_custom_ui_sso_sign_in_handler,
1019 )
1021 microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
1022 google_client_id: Final = os.getenv("GOOGLE_CLIENT_ID", None)
1023 generic_client_id: Final = os.getenv("GENERIC_CLIENT_ID", None)
1025 ####### Check if UI is disabled #######
1026 _disable_ui_flag: Final = os.getenv("DISABLE_ADMIN_UI")
1027 if _disable_ui_flag is not None:
1028 is_disabled: Final = str_to_bool(value=_disable_ui_flag)
1029 if is_disabled:
1030 return admin_ui_disabled()
1032 ####### Check if user is a Enterprise / Premium User #######
1033 if (
1034 microsoft_client_id is not None
1035 or google_client_id is not None
1036 or generic_client_id is not None
1037 or SAMLAuthHandler.is_saml_configured()
1038 ):
1039 await _raise_if_sso_exceeds_free_user_limit(premium_user, prisma_client)
1041 ####### Detect DB + MASTER KEY in .env #######
1042 missing_env_vars: Final = show_missing_vars_in_env()
1043 if missing_env_vars is not None:
1044 return missing_env_vars
1046 # get url from request - always use regular callback, but set state for CLI
1047 redirect_url: Final = SSOAuthenticationHandler.get_redirect_url_for_sso(
1048 request=request,
1049 sso_callback_route="sso/callback",
1050 )
1052 if source == LITELLM_CLI_SOURCE_IDENTIFIER:
1053 _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache)
1055 # Store CLI login handle in state for OAuth flow
1056 cli_state: Final[str | None] = SSOAuthenticationHandler._get_cli_state(
1057 source=source,
1058 key=key,
1059 user_code=(user_code if _cli_sso_verification_uri_complete_enabled() else None),
1060 )
1062 # check if user defined a custom auth sso sign in handler, if yes, use it
1063 if user_custom_ui_sso_sign_in_handler is not None:
1064 try:
1065 from litellm_enterprise.proxy.auth.custom_sso_handler import (
1066 EnterpriseCustomSSOHandler,
1067 )
1069 return await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in(
1070 request=request,
1071 )
1072 except ImportError:
1073 raise ValueError(
1074 "Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise."
1075 )
1077 if (
1078 microsoft_client_id is None
1079 and google_client_id is None
1080 and generic_client_id is None
1081 and SAMLAuthHandler.is_saml_configured()
1082 ):
1083 verbose_proxy_logger.info("Redirecting to SAML SSO login")
1084 return await SAMLAuthHandler.build_login_redirect(
1085 request=request,
1086 cache=user_api_key_cache,
1087 relay_state=return_to,
1088 )
1090 # Check if we should use SSO handler
1091 if (
1092 SSOAuthenticationHandler.should_use_sso_handler(
1093 microsoft_client_id=microsoft_client_id,
1094 google_client_id=google_client_id,
1095 generic_client_id=generic_client_id,
1096 )
1097 is True
1098 ):
1099 verbose_proxy_logger.info("Redirecting to SSO login for %s", redirect_url)
1100 sso_redirect: Final = await SSOAuthenticationHandler.get_sso_login_redirect(
1101 redirect_url=redirect_url,
1102 microsoft_client_id=microsoft_client_id,
1103 google_client_id=google_client_id,
1104 generic_client_id=generic_client_id,
1105 state=cli_state,
1106 request=request,
1107 )
1108 if sso_redirect is not None:
1109 _persist_return_to_cookie(sso_redirect, return_to, request)
1110 return sso_redirect
1112 from fastapi.responses import HTMLResponse
1114 hide_default_credentials_hint: Final = should_hide_default_credentials_hint(general_settings)
1115 form_response: Final = HTMLResponse(
1116 content=build_ui_login_form(
1117 show_deprecation_banner=True,
1118 hide_default_credentials_hint=hide_default_credentials_hint,
1119 ),
1120 status_code=200,
1121 )
1122 # Preserve return_to across the username/password sign-in too, via the SAME shared, never-raising
1123 # helper the SSO branch uses, so /login can resume the connect flow instead of dead-ending at the
1124 # dashboard. One implementation → the two sign-in branches cannot diverge (and the login form always
1125 # renders, since the helper never raises on a bad return_to).
1126 _persist_return_to_cookie(form_response, return_to, request)
1127 return form_response
1130def generic_response_convertor(
1131 response,
1132 jwt_handler: JWTHandler,
1133 sso_jwt_handler: JWTHandler | None = None,
1134 role_mappings: Optional["RoleMappings"] = None,
1135 team_mappings: Optional["TeamMappings"] = None,
1136) -> CustomOpenID:
1137 generic_user_id_attribute_name: Final = os.getenv("GENERIC_USER_ID_ATTRIBUTE", "preferred_username")
1138 generic_user_display_name_attribute_name: Final = os.getenv("GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "sub")
1139 generic_user_email_attribute_name: Final = os.getenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email")
1141 generic_user_first_name_attribute_name: Final = os.getenv("GENERIC_USER_FIRST_NAME_ATTRIBUTE", "first_name")
1142 generic_user_last_name_attribute_name: Final = os.getenv("GENERIC_USER_LAST_NAME_ATTRIBUTE", "last_name")
1144 generic_provider_attribute_name: Final = os.getenv("GENERIC_USER_PROVIDER_ATTRIBUTE", "provider")
1146 generic_user_role_attribute_name: Final = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role")
1148 generic_user_extra_attributes: Final = os.getenv("GENERIC_USER_EXTRA_ATTRIBUTES", None)
1150 verbose_proxy_logger.debug(
1151 " generic_user_id_attribute_name: %s\n generic_user_email_attribute_name: %s",
1152 generic_user_id_attribute_name,
1153 generic_user_email_attribute_name,
1154 )
1156 all_teams: Final = []
1157 if sso_jwt_handler is not None:
1158 team_ids = sso_jwt_handler.get_all_jwt_team_ids(cast(dict, response))
1159 all_teams.extend(team_ids)
1161 if team_mappings is not None and team_mappings.team_ids_jwt_field is not None:
1162 team_ids_from_db_mapping: Final[list[str] | None] = get_nested_value(
1163 data=cast(dict, response),
1164 key_path=team_mappings.team_ids_jwt_field,
1165 default=[],
1166 )
1167 if team_ids_from_db_mapping:
1168 all_teams.extend(team_ids_from_db_mapping)
1169 verbose_proxy_logger.debug(
1170 "Loaded team_ids from DB team_mappings.team_ids_jwt_field='%s': %s",
1171 team_mappings.team_ids_jwt_field,
1172 team_ids_from_db_mapping,
1173 )
1174 else:
1175 team_ids = jwt_handler.get_all_jwt_team_ids(cast(dict, response))
1176 all_teams.extend(team_ids)
1178 # Determine user role based on role_mappings if available
1179 # Only apply role_mappings for GENERIC SSO provider
1180 user_role: LitellmUserRoles | None = None
1182 if role_mappings is not None and role_mappings.provider.lower() in [
1183 "generic",
1184 "okta",
1185 ]:
1186 # Use role_mappings to determine role from groups
1187 group_claim: Final = role_mappings.group_claim
1188 user_groups_raw: Final[object] = get_nested_value(response, group_claim)
1190 # Handle different formats: could be a list, string (comma-separated), or single value
1191 user_groups: list[str] = []
1192 if isinstance(user_groups_raw, list):
1193 raw_groups: Final[Sequence[object]] = user_groups_raw
1194 user_groups = [str(g) for g in raw_groups]
1195 elif isinstance(user_groups_raw, str):
1196 # Handle comma-separated string
1197 user_groups = [g.strip() for g in user_groups_raw.split(",") if g.strip()]
1198 elif user_groups_raw is not None:
1199 # Single value
1200 user_groups = [str(user_groups_raw)]
1202 if user_groups:
1203 user_role = determine_role_from_groups(user_groups, role_mappings)
1204 verbose_proxy_logger.debug(
1205 "Determined role '%s' from groups '%s' using role_mappings",
1206 user_role.value if user_role else None,
1207 user_groups,
1208 )
1209 else:
1210 # No groups found, use default_role
1211 user_role = role_mappings.default_role
1212 verbose_proxy_logger.debug(
1213 "No groups found in '%s', using default_role: %s", group_claim, role_mappings.default_role
1214 )
1216 # Fallback to existing logic if role_mappings not used
1217 if user_role is None:
1218 user_role_from_sso: Final = get_nested_value(response, generic_user_role_attribute_name)
1219 if user_role_from_sso is not None:
1220 role: Final = get_litellm_user_role(user_role_from_sso)
1221 if role is not None:
1222 user_role = role
1223 verbose_proxy_logger.debug(
1224 "Found valid LitellmUserRoles '%s' from SSO attribute '%s'",
1225 role.value,
1226 generic_user_role_attribute_name,
1227 )
1229 # Build extra_fields dict from GENERIC_USER_EXTRA_ATTRIBUTES if specified
1230 extra_fields: dict[str, object] | None = None
1231 if generic_user_extra_attributes:
1232 extra_fields = {}
1233 for attr_name in generic_user_extra_attributes.split(","):
1234 attr_name = attr_name.strip()
1235 extra_fields[attr_name] = get_nested_value(response, attr_name)
1237 return CustomOpenID(
1238 id=get_nested_value(response, generic_user_id_attribute_name),
1239 display_name=get_nested_value(response, generic_user_display_name_attribute_name),
1240 email=normalize_email(get_nested_value(response, generic_user_email_attribute_name)),
1241 first_name=get_nested_value(response, generic_user_first_name_attribute_name),
1242 last_name=get_nested_value(response, generic_user_last_name_attribute_name),
1243 provider=get_nested_value(response, generic_provider_attribute_name),
1244 team_ids=all_teams,
1245 user_role=user_role,
1246 extra_fields=extra_fields,
1247 )
1250def _setup_generic_sso_env_vars(
1251 generic_client_id: str, redirect_url: str
1252) -> tuple[str, list[str], str, str, str, bool]:
1253 """Setup and validate Generic SSO environment variables."""
1254 generic_client_secret: Final = os.getenv("GENERIC_CLIENT_SECRET", None)
1255 generic_scope: Final = os.getenv("GENERIC_SCOPE", "openid email profile").split(" ")
1256 generic_authorization_endpoint: Final = os.getenv("GENERIC_AUTHORIZATION_ENDPOINT", None)
1257 generic_token_endpoint: Final = os.getenv("GENERIC_TOKEN_ENDPOINT", None)
1258 generic_userinfo_endpoint: Final = os.getenv("GENERIC_USERINFO_ENDPOINT", None)
1259 generic_include_client_id: Final = os.getenv("GENERIC_INCLUDE_CLIENT_ID", "false").lower() == "true"
1261 # Validate required environment variables
1262 if generic_client_secret is None:
1263 raise ProxyException(
1264 message="GENERIC_CLIENT_SECRET not set. Set it in .env file",
1265 type=ProxyErrorTypes.auth_error,
1266 param="GENERIC_CLIENT_SECRET",
1267 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
1268 )
1269 if generic_authorization_endpoint is None:
1270 raise ProxyException(
1271 message="GENERIC_AUTHORIZATION_ENDPOINT not set. Set it in .env file",
1272 type=ProxyErrorTypes.auth_error,
1273 param="GENERIC_AUTHORIZATION_ENDPOINT",
1274 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
1275 )
1276 if generic_token_endpoint is None:
1277 raise ProxyException(
1278 message="GENERIC_TOKEN_ENDPOINT not set. Set it in .env file",
1279 type=ProxyErrorTypes.auth_error,
1280 param="GENERIC_TOKEN_ENDPOINT",
1281 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
1282 )
1283 if generic_userinfo_endpoint is None:
1284 raise ProxyException(
1285 message="GENERIC_USERINFO_ENDPOINT not set. Set it in .env file",
1286 type=ProxyErrorTypes.auth_error,
1287 param="GENERIC_USERINFO_ENDPOINT",
1288 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
1289 )
1291 verbose_proxy_logger.debug(
1292 "authorization_endpoint: %s\ntoken_endpoint: %s\nuserinfo_endpoint: %s",
1293 generic_authorization_endpoint,
1294 generic_token_endpoint,
1295 generic_userinfo_endpoint,
1296 )
1297 verbose_proxy_logger.debug("GENERIC_REDIRECT_URI: %s\nGENERIC_CLIENT_ID: %s\n", redirect_url, generic_client_id)
1299 return (
1300 generic_client_secret,
1301 generic_scope,
1302 generic_authorization_endpoint,
1303 generic_token_endpoint,
1304 generic_userinfo_endpoint,
1305 generic_include_client_id,
1306 )
1309async def _setup_team_mappings() -> Optional["TeamMappings"]:
1310 """Setup team mappings from SSO database settings."""
1311 from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings
1313 team_mappings: TeamMappings | None = None
1314 try:
1315 from litellm.proxy.utils import get_prisma_client_or_throw
1317 prisma_client: Final = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy")
1319 sso_db_record: Final = await _sso_config_db(SSOConfigRepository(prisma_client)).find_unique(
1320 where={"id": "sso_config"}
1321 )
1323 if sso_db_record and sso_db_record.sso_settings:
1324 sso_settings_dict: Final = dict(sso_db_record.sso_settings)
1325 team_mappings_data: Final = sso_settings_dict.get("team_mappings")
1327 if team_mappings_data:
1328 if isinstance(team_mappings_data, dict):
1329 team_mappings = TeamMappings(**team_mappings_data)
1330 elif isinstance(team_mappings_data, TeamMappings):
1331 team_mappings = team_mappings_data
1333 if team_mappings and team_mappings.team_ids_jwt_field:
1334 verbose_proxy_logger.debug(
1335 "Loaded team_mappings with team_ids_jwt_field: '%s'", team_mappings.team_ids_jwt_field
1336 )
1337 except Exception as e:
1338 verbose_proxy_logger.debug(
1339 "Could not load team_mappings from database: %s. Continuing with config-based team mapping.", e
1340 )
1342 return team_mappings
1345async def _setup_role_mappings() -> Optional["RoleMappings"]:
1346 """Setup role mappings from SSO database settings."""
1347 role_mappings: RoleMappings | None = None
1348 try:
1349 from litellm.proxy.utils import get_prisma_client_or_throw
1351 prisma_client: Final = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy")
1353 sso_db_record: Final = await _sso_config_db(SSOConfigRepository(prisma_client)).find_unique(
1354 where={"id": "sso_config"}
1355 )
1357 if sso_db_record and sso_db_record.sso_settings:
1358 sso_settings_dict: Final = dict(sso_db_record.sso_settings)
1359 role_mappings_data = sso_settings_dict.get("role_mappings")
1361 if role_mappings_data:
1362 if isinstance(role_mappings_data, dict):
1363 role_mappings = RoleMappings(**role_mappings_data)
1364 elif isinstance(role_mappings_data, RoleMappings):
1365 role_mappings = role_mappings_data
1367 if role_mappings:
1368 verbose_proxy_logger.debug("Loaded role_mappings for provider '%s'", role_mappings.provider)
1369 except Exception as e:
1370 verbose_proxy_logger.debug(
1371 "Could not load role_mappings from database: %s. Continuing with existing role logic.", e
1372 )
1374 generic_role_mappings: Final = os.getenv("GENERIC_ROLE_MAPPINGS_ROLES", None)
1375 generic_role_mappings_group_claim: Final = os.getenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None)
1376 generic_role_mappings_default_role: Final = os.getenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", None)
1377 if generic_role_mappings is not None:
1378 verbose_proxy_logger.debug("Found role_mappings for generic provider in environment variables")
1379 import ast
1381 try:
1382 generic_user_role_mappings_data: dict[LitellmUserRoles, list[str]] = ast.literal_eval(generic_role_mappings)
1383 if isinstance(generic_user_role_mappings_data, dict):
1384 role_mappings_data = {
1385 "provider": "generic",
1386 "group_claim": generic_role_mappings_group_claim,
1387 "default_role": generic_role_mappings_default_role,
1388 "roles": generic_user_role_mappings_data,
1389 }
1391 role_mappings = RoleMappings(**role_mappings_data)
1392 verbose_proxy_logger.debug(
1393 "Loaded role_mappings from environments for provider '%s'.", role_mappings.provider
1394 )
1395 return role_mappings
1396 except TypeError as e:
1397 verbose_proxy_logger.warning(
1398 "Error decoding role mappings from environment variables: %s. Continuing with existing role logic.", e
1399 )
1400 return role_mappings
1403def _parse_generic_sso_headers() -> dict[str, str]:
1404 """Parse comma-separated GENERIC_SSO_HEADERS env var into a dict."""
1405 raw: Final = os.getenv("GENERIC_SSO_HEADERS", None)
1406 if raw is None:
1407 return {}
1408 result: Final[dict[str, str]] = {}
1409 for header in raw.split(","):
1410 header = header.strip()
1411 if header:
1412 key, value = header.split("=")
1413 result[key] = value
1414 return result
1417def _handle_generic_sso_error(
1418 e: Exception,
1419 generic_authorization_endpoint: str | None,
1420 generic_token_endpoint: str | None,
1421 additional_headers: dict,
1422) -> NoReturn:
1423 """Handle errors from generic SSO verify_and_process. Always re-raises."""
1424 error_message: Final = str(e)
1426 # Surface a helpful PKCE misconfiguration hint only when:
1427 # 1. The error mentions PKCE/code verifier, AND
1428 # 2. PKCE is not currently configured (GENERIC_CLIENT_USE_PKCE != true)
1429 pkce_configured: Final = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
1430 if not pkce_configured and ("PKCE" in error_message or "code verifier" in error_message.lower()):
1431 is_okta: Final = (generic_authorization_endpoint and "okta" in generic_authorization_endpoint.lower()) or (
1432 generic_token_endpoint and "okta" in generic_token_endpoint.lower()
1433 )
1434 provider_name: Final = "Okta" if is_okta else "Your OAuth provider"
1436 detailed_message = (
1437 f"SSO authentication failed: {provider_name} requires PKCE (Proof Key for Code Exchange) "
1438 f"but it's not enabled in your LiteLLM configuration.\n\n"
1439 f"SOLUTION: Add this environment variable and restart your proxy:\n"
1440 f" GENERIC_CLIENT_USE_PKCE=true\n\n"
1441 )
1442 if is_okta:
1443 detailed_message += (
1444 "For AWS ECS: Add the environment variable to your task definition.\n"
1445 "For Docker: Add -e GENERIC_CLIENT_USE_PKCE=true to your docker run command.\n"
1446 "For .env file: Add GENERIC_CLIENT_USE_PKCE=true to your .env file.\n\n"
1447 )
1448 detailed_message += f"Original error: {error_message}"
1450 raise ProxyException(
1451 message=detailed_message,
1452 type=ProxyErrorTypes.auth_error,
1453 param="GENERIC_CLIENT_USE_PKCE",
1454 code=status.HTTP_401_UNAUTHORIZED,
1455 )
1457 if isinstance(e, ProxyException):
1458 verbose_proxy_logger.error(
1459 "SSO authentication failed: %s. Passed in headers: %s",
1460 e,
1461 additional_headers,
1462 )
1463 else:
1464 verbose_proxy_logger.exception(
1465 "Error verifying and processing generic SSO: %s. Passed in headers: %s",
1466 e,
1467 additional_headers,
1468 )
1469 raise e
1472async def get_generic_sso_response(
1473 request: Request,
1474 jwt_handler: JWTHandler,
1475 sso_jwt_handler: JWTHandler | None, # sso specific jwt handler - used for restricted sso group access control
1476 generic_client_id: str,
1477 redirect_url: str,
1478) -> tuple[
1479 OpenID | dict, dict | None, dict | None, SSOIdentityAssertion | None
1480]: # (result, received_response, access_token_payload, sso_assertion)
1481 # make generic sso provider
1482 from fastapi_sso.sso.base import DiscoveryDocument
1483 from fastapi_sso.sso.generic import create_provider
1485 received_response: dict | None = None
1486 sso_assertion: SSOIdentityAssertion | None = None
1488 # Setup environment variables
1489 (
1490 generic_client_secret,
1491 generic_scope,
1492 generic_authorization_endpoint,
1493 generic_token_endpoint,
1494 generic_userinfo_endpoint,
1495 generic_include_client_id,
1496 ) = _setup_generic_sso_env_vars(generic_client_id, redirect_url)
1498 discovery: Final = DiscoveryDocument(
1499 authorization_endpoint=generic_authorization_endpoint,
1500 token_endpoint=generic_token_endpoint,
1501 userinfo_endpoint=generic_userinfo_endpoint,
1502 )
1504 role_mappings: Final = await _setup_role_mappings()
1505 team_mappings: Final = await _setup_team_mappings()
1506 generic_include_token_claims: Final = os.getenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "false").lower() == "true"
1508 def response_convertor(response: Mapping[str, object], httpx_session: object):
1509 nonlocal received_response # return for user debugging
1510 response_id_token: Final = response.get("id_token")
1511 response_access_token: Final = response.get("access_token")
1512 id_token: Final = (
1513 response_id_token if isinstance(response_id_token, str) and response_id_token else generic_sso.id_token
1514 )
1515 access_token: Final = (
1516 response_access_token
1517 if isinstance(response_access_token, str) and response_access_token
1518 else generic_sso.access_token
1519 )
1520 claims: Final = (
1521 _merge_sso_token_claims(
1522 userinfo=response,
1523 id_token=id_token,
1524 access_token=access_token,
1525 )
1526 if generic_include_token_claims
1527 else response
1528 )
1529 received_response = { # mutable-ok: preserve the existing dict return contract
1530 key: value for key, value in claims.items() if key not in _OAUTH_TOKEN_FIELDS
1531 }
1532 return generic_response_convertor(
1533 response=claims,
1534 jwt_handler=jwt_handler,
1535 sso_jwt_handler=sso_jwt_handler,
1536 role_mappings=role_mappings,
1537 team_mappings=team_mappings,
1538 )
1540 SSOProvider: Final = create_provider(
1541 name="oidc",
1542 discovery_document=discovery,
1543 response_convertor=response_convertor,
1544 )
1545 generic_sso: Final = SSOProvider(
1546 client_id=generic_client_id,
1547 client_secret=generic_client_secret,
1548 redirect_uri=redirect_url,
1549 allow_insecure_http=True,
1550 scope=generic_scope,
1551 )
1552 verbose_proxy_logger.debug("calling generic_sso.verify_and_process")
1553 additional_generic_sso_headers_dict: Final = _parse_generic_sso_headers()
1555 code_verifier: str | None = None # assigned inside try; initialized for type tracking
1556 access_token_payload: dict | None = None # decoded JWT access token claims
1558 try:
1559 token_exchange_params: Final = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
1560 request=request,
1561 generic_include_client_id=generic_include_client_id,
1562 )
1564 # Extract code_verifier (and the cache key for deferred deletion) before calling fastapi-sso
1565 code_verifier = token_exchange_params.pop("code_verifier", None)
1566 pkce_cache_key: Final = token_exchange_params.pop("_pkce_cache_key", None)
1568 # Get authorization code from query params (only used in the PKCE path below;
1569 # the non-PKCE path delegates to verify_and_process which handles OAuth error
1570 # callbacks — user-denied, CSRF mismatch — internally).
1571 authorization_code: Final = request.query_params.get("code")
1573 if code_verifier:
1574 # State-to-session-cookie binding. The non-PKCE branch below
1575 # delegates to fastapi-sso's ``verify_and_process``, which
1576 # performs its own session-cookie check. The PKCE branch
1577 # bypasses that helper, so we validate the URL ``state``
1578 # against the ``litellm_oauth_state`` cookie set on the
1579 # redirect response — without this an attacker can pre-mint
1580 # a state + cached PKCE verifier and hijack a victim's auth
1581 # code (Login-CSRF / token theft).
1582 url_state: Final = request.query_params.get("state")
1583 cookie_state: Final = request.cookies.get("litellm_oauth_state")
1584 if not url_state or not cookie_state or not secrets.compare_digest(url_state, cookie_state):
1585 raise ProxyException(
1586 message=("Invalid OAuth state parameter — does not match the browser-bound state cookie."),
1587 type=ProxyErrorTypes.auth_error,
1588 param="state",
1589 code=status.HTTP_400_BAD_REQUEST,
1590 )
1591 if not authorization_code:
1592 raise ProxyException(
1593 message="Missing authorization code in callback",
1594 type=ProxyErrorTypes.auth_error,
1595 param="code",
1596 code=status.HTTP_400_BAD_REQUEST,
1597 )
1598 if not generic_client_id:
1599 raise ProxyException(
1600 message="GENERIC_CLIENT_ID must be set when PKCE is enabled",
1601 type=ProxyErrorTypes.auth_error,
1602 param="GENERIC_CLIENT_ID",
1603 code=status.HTTP_401_UNAUTHORIZED,
1604 )
1605 if not generic_token_endpoint:
1606 raise ProxyException(
1607 message="GENERIC_TOKEN_ENDPOINT must be set when PKCE is enabled",
1608 type=ProxyErrorTypes.auth_error,
1609 param="GENERIC_TOKEN_ENDPOINT",
1610 code=status.HTTP_401_UNAUTHORIZED,
1611 )
1612 # All guards above raise, so authorization_code is a non-empty str here.
1613 # Use an explicit type guard rather than assert (assert is a no-op with -O).
1614 if not isinstance(authorization_code, str):
1615 raise ProxyException(
1616 message="Missing authorization code in callback",
1617 type=ProxyErrorTypes.auth_error,
1618 param="code",
1619 code=status.HTTP_400_BAD_REQUEST,
1620 )
1621 combined_response: Final = await SSOAuthenticationHandler._pkce_token_exchange(
1622 authorization_code=authorization_code,
1623 code_verifier=code_verifier,
1624 client_id=generic_client_id,
1625 client_secret=generic_client_secret,
1626 token_endpoint=generic_token_endpoint,
1627 userinfo_endpoint=generic_userinfo_endpoint,
1628 include_client_id=generic_include_client_id,
1629 redirect_url=redirect_url,
1630 additional_headers=additional_generic_sso_headers_dict,
1631 )
1632 # Pass the full response so custom response_convertor implementations
1633 # can access all fields (including id_token for claim extraction).
1634 result = response_convertor(combined_response, generic_sso)
1635 sso_assertion = assertion_from_sso_login(
1636 combined_response.get("id_token"), combined_response.get("refresh_token")
1637 )
1638 # In the PKCE path verify_and_process is skipped, so generic_sso.access_token
1639 # is never set. Read the token directly from the exchange response instead so
1640 # process_sso_jwt_access_token can extract JWT-embedded roles/teams.
1641 access_token_str: str | None = combined_response.get("access_token")
1642 else:
1643 result = await generic_sso.verify_and_process(
1644 request,
1645 params=token_exchange_params,
1646 headers=additional_generic_sso_headers_dict,
1647 )
1648 access_token_str = generic_sso.access_token
1649 sso_assertion = assertion_from_sso_login(generic_sso.id_token, generic_sso.refresh_token)
1651 access_token_payload = process_sso_jwt_access_token(
1652 access_token_str, sso_jwt_handler, result, role_mappings=role_mappings
1653 )
1654 # Delete the single-use PKCE verifier only after all downstream processing
1655 # (response_convertor and process_sso_jwt_access_token) has completed
1656 # successfully. Deleting earlier would consume the verifier on a transient
1657 # failure, forcing the user to restart the entire OAuth flow from scratch.
1658 if pkce_cache_key:
1659 await SSOAuthenticationHandler._delete_pkce_verifier(pkce_cache_key)
1661 except Exception as e:
1662 _handle_generic_sso_error(
1663 e,
1664 generic_authorization_endpoint,
1665 generic_token_endpoint,
1666 additional_generic_sso_headers_dict,
1667 )
1668 verbose_proxy_logger.debug("generic result: %s", result)
1669 return result or {}, received_response, access_token_payload, sso_assertion
1672RetentionCheck: TypeAlias = Callable[[], Awaitable[bool]] # mutable-ok: Callable parameter syntax
1675async def warn_if_id_jag_assertion_uncaptured(
1676 assertion: SSOIdentityAssertion | None, *, retention_enabled: RetentionCheck | None = None
1677) -> None:
1678 """Say, at the one moment it is knowable, that this login gave an ``oauth2_id_jag`` server
1679 nothing to spend. Without it the operator only ever sees the per-request failure, which cannot
1680 tell a user who has never signed in from a provider that will never capture. Kept strictly
1681 diagnostic: a store outage is swallowed, since a login must not fail over a log line."""
1682 if assertion is not None:
1683 return
1684 try:
1685 check: Final = retention_enabled if retention_enabled is not None else ema_assertion_retention_enabled
1686 if not await check():
1687 return
1688 except Exception as exc: # noqa: BLE001 # diagnostics must never break the login
1689 verbose_proxy_logger.debug("Could not check for oauth2_id_jag MCP servers after SSO login: %s", exc)
1690 return
1691 gap: Final = id_jag_assertion_capture_gap()
1692 verbose_proxy_logger.warning(
1693 "SSO login captured no IdP identity assertion while an oauth2_id_jag MCP server is registered: %s",
1694 gap if gap is not None else "the identity provider's token response carried no usable id_token",
1695 )
1698async def warn_if_id_jag_capture_gap(*, retention_enabled: RetentionCheck | None = None) -> None:
1699 gap: Final = id_jag_assertion_capture_gap()
1700 if gap is None:
1701 return
1702 try:
1703 check: Final = retention_enabled if retention_enabled is not None else ema_assertion_retention_enabled
1704 if not await check():
1705 return
1706 except Exception as exc: # noqa: BLE001 # diagnostics must never break the page they annotate
1707 verbose_proxy_logger.debug("Could not check for oauth2_id_jag MCP servers: %s", exc)
1708 return
1709 verbose_proxy_logger.warning("SSO debug callback ran with an oauth2_id_jag capture gap: %s", gap)
1712async def create_team_member_add_task(team_id, user_info):
1713 """Create a task for adding a member to a team."""
1714 try:
1715 member: Final = Member(user_id=user_info.user_id, role="user")
1716 team_member_add_request: Final = TeamMemberAddRequest(
1717 member=member,
1718 team_id=team_id,
1719 )
1720 return await team_member_add(
1721 data=team_member_add_request,
1722 user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
1723 )
1724 except Exception as e:
1725 verbose_proxy_logger.debug("[Non-Blocking] Error trying to add sso user to db: %s", e)
1728async def add_missing_team_member(user_info: NewUserResponse | LiteLLM_UserTable, sso_teams: list[str]):
1729 """
1730 - Get missing teams (diff b/w user_info.team_ids and sso_teams)
1731 - Add missing user to missing teams
1732 """
1733 # Handle None as empty list for new users
1734 user_teams: Final = user_info.teams if user_info.teams is not None else []
1735 missing_teams: Final = set(sso_teams) - set(user_teams)
1736 missing_teams_list: Final = list(missing_teams)
1737 tasks = []
1738 tasks = [create_team_member_add_task(team_id, user_info) for team_id in missing_teams_list]
1740 try:
1741 await asyncio.gather(*tasks)
1742 except Exception as e:
1743 verbose_proxy_logger.debug("[Non-Blocking] Error trying to add sso user to db: %s", e)
1746def get_disabled_non_admin_personal_key_creation():
1747 key_generation_settings: Final = litellm.key_generation_settings
1748 if key_generation_settings is None:
1749 return False
1750 personal_key_generation: Final = key_generation_settings.get("personal_key_generation") or {}
1751 allowed_user_roles: Final = personal_key_generation.get("allowed_user_roles") or []
1752 return bool("proxy_admin" in allowed_user_roles)
1755async def get_existing_user_info_from_db(
1756 user_id: str | None,
1757 user_email: str | None,
1758 prisma_client: PrismaClient,
1759 user_api_key_cache: UserApiKeyCache,
1760 proxy_logging_obj: ProxyLogging,
1761) -> LiteLLM_UserTable | None:
1762 try:
1763 user_info = await get_user_object(
1764 user_id=user_id,
1765 user_email=user_email,
1766 prisma_client=prisma_client,
1767 user_api_key_cache=user_api_key_cache,
1768 user_id_upsert=False,
1769 parent_otel_span=None,
1770 proxy_logging_obj=proxy_logging_obj,
1771 sso_user_id=user_id,
1772 )
1773 except Exception as e:
1774 verbose_proxy_logger.debug("Error getting user object: %s", e)
1775 user_info = None
1777 return user_info
1780async def get_user_info_from_db(
1781 result: CustomOpenID | OpenID | dict,
1782 prisma_client: PrismaClient,
1783 user_api_key_cache: UserApiKeyCache,
1784 proxy_logging_obj: ProxyLogging,
1785 user_email: str | None,
1786 user_defined_values: SSOUserDefinedValues | None,
1787 alternate_user_id: str | None = None,
1788) -> LiteLLM_UserTable | NewUserResponse | None:
1789 try:
1790 potential_user_ids: Final = []
1791 if alternate_user_id is not None:
1792 potential_user_ids.append(alternate_user_id)
1793 if not isinstance(result, dict):
1794 _id = getattr(result, "id", None)
1795 if _id is not None and isinstance(_id, str):
1796 potential_user_ids.append(_id)
1797 else:
1798 _id = result.get("id", None)
1799 if _id is not None and isinstance(_id, str):
1800 potential_user_ids.append(_id)
1802 user_email = normalize_email(
1803 getattr(result, "email", None) if not isinstance(result, dict) else result.get("email", None)
1804 )
1806 user_info: LiteLLM_UserTable | NewUserResponse | None = None
1808 for user_id in potential_user_ids:
1809 user_info = await get_existing_user_info_from_db(
1810 user_id=user_id,
1811 user_email=user_email,
1812 prisma_client=prisma_client,
1813 user_api_key_cache=user_api_key_cache,
1814 proxy_logging_obj=proxy_logging_obj,
1815 )
1816 if user_info is not None:
1817 break
1819 verbose_proxy_logger.debug(
1820 "user_info: %s; litellm.default_internal_user_params: %s", user_info, litellm.default_internal_user_params
1821 )
1823 # Upsert SSO User to LiteLLM DB
1824 user_info = await SSOAuthenticationHandler.upsert_sso_user(
1825 result=result,
1826 user_info=user_info,
1827 user_email=user_email,
1828 user_defined_values=user_defined_values,
1829 prisma_client=prisma_client,
1830 )
1832 await SSOAuthenticationHandler.add_user_to_teams_from_sso_response(
1833 result=result,
1834 user_info=user_info,
1835 )
1837 return user_info
1838 except Exception as e:
1839 verbose_proxy_logger.exception("[Non-Blocking] Error trying to add sso user to db: %s", e)
1841 return None
1844def _should_use_role_from_sso_response(sso_role: str | None) -> bool:
1845 """returns true if SSO upsert should use the 'role' defined on the SSO response"""
1846 if sso_role is None:
1847 return False
1849 if not is_valid_litellm_user_role(sso_role):
1850 verbose_proxy_logger.debug(
1851 "SSO role '%s' is not a valid LiteLLM user role. Ignoring role from SSO response. See LitellmUserRoles enum for valid roles.",
1852 sso_role,
1853 )
1854 return False
1855 return True
1858def _build_sso_user_update_data(
1859 result: Union["CustomOpenID", OpenID, dict] | None,
1860 user_email: str | None,
1861 user_id: str | None,
1862) -> dict[str, object]:
1863 """
1864 Build the update data dictionary for SSO user upsert.
1866 Args:
1867 result: The SSO response containing user information
1868 user_email: The user's email from SSO
1869 user_id: The user's ID for logging purposes
1871 Returns:
1872 dict: Update data containing user_email and optionally user_role if valid
1873 """
1874 update_data: Final[dict[str, object]] = {"user_email": normalize_email(user_email)}
1876 # Get SSO role from result and include if valid
1877 sso_role: Final = getattr(result, "user_role", None)
1878 if sso_role is not None:
1879 # Convert enum to string if needed
1880 sso_role_str: Final = sso_role.value if isinstance(sso_role, LitellmUserRoles) else sso_role
1882 # Only include if it's a valid LiteLLM role
1883 if _should_use_role_from_sso_response(sso_role_str):
1884 update_data["user_role"] = sso_role_str
1885 verbose_proxy_logger.info("Updating user %s role from SSO: %s", user_id, sso_role_str)
1887 return update_data
1890async def _sync_user_role_from_jwt_role_map(
1891 jwt_handler: JWTHandler | None,
1892 received_response: dict | None,
1893 user_info: LiteLLM_UserTable | NewUserResponse | None,
1894 prisma_client: PrismaClient,
1895 user_api_key_cache: UserApiKeyCache,
1896 user_defined_values: SSOUserDefinedValues | None,
1897) -> None:
1898 """
1899 Apply jwt_litellm_role_map during SSO login.
1901 When jwt_litellm_role_map is configured with sync_user_role_and_teams=True,
1902 this ensures SSO users get the same role mapping as API/JWT users. Without
1903 this, the SSO path falls back to INTERNAL_USER_VIEW_ONLY for roles that
1904 don't directly match LitellmUserRoles enum values.
1905 """
1906 if jwt_handler is None or received_response is None:
1907 return
1908 if not jwt_handler.litellm_jwtauth.sync_user_role_and_teams:
1909 return
1910 if not jwt_handler.litellm_jwtauth.jwt_litellm_role_map:
1911 return
1913 mapped_role: Final = jwt_handler.map_jwt_role_to_litellm_role(received_response)
1914 if mapped_role is None:
1915 return
1917 verbose_proxy_logger.info("SSO jwt_litellm_role_map matched role: %s", mapped_role.value)
1919 # Update user_defined_values so downstream code uses the mapped role
1920 if user_defined_values is not None:
1921 user_defined_values["user_role"] = mapped_role.value
1923 # Update existing DB record if role differs
1924 if user_info is not None and user_info.user_role != mapped_role.value:
1925 await _user_meta_db(UserRepository(prisma_client)).update(
1926 where={"user_id": user_info.user_id},
1927 data={"user_role": mapped_role.value},
1928 )
1929 user_info.user_role = mapped_role.value
1930 await user_api_key_cache.async_set_cache(
1931 key=user_info.user_id,
1932 value=user_info,
1933 model_type=LiteLLM_UserTable,
1934 )
1937def apply_user_info_values_to_sso_user_defined_values(
1938 user_info: LiteLLM_UserTable | NewUserResponse | None,
1939 user_defined_values: SSOUserDefinedValues | None,
1940) -> SSOUserDefinedValues | None:
1941 if user_defined_values is None:
1942 return None
1943 if user_info is not None and user_info.user_id is not None:
1944 user_defined_values["user_id"] = user_info.user_id
1946 # SSO role takes precedence - only use DB role if SSO didn't provide one
1947 # This ensures SSO is the authoritative source for user roles
1948 sso_role: Final = user_defined_values.get("user_role")
1949 db_role: Final = user_info.user_role if user_info else None
1951 if _should_use_role_from_sso_response(sso_role):
1952 # SSO provided a valid role, keep it and log that we're using it
1953 verbose_proxy_logger.info("Using SSO role: %s (DB role was: %s)", sso_role, db_role)
1954 else:
1955 # SSO didn't provide a valid role, fall back to DB role or default
1956 if user_info is None or user_info.user_role is None:
1957 user_defined_values["user_role"] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value
1958 verbose_proxy_logger.debug("No SSO or DB role found, using default: INTERNAL_USER_VIEW_ONLY")
1959 else:
1960 user_defined_values["user_role"] = user_info.user_role
1961 verbose_proxy_logger.debug("Using DB role: %s", user_info.user_role)
1963 # Preserve the user's existing models from the database
1964 if user_info is not None and hasattr(user_info, "models") and user_info.models:
1965 user_defined_values["models"] = user_info.models
1967 return user_defined_values
1970async def check_and_update_if_proxy_admin_id(user_role: str, user_id: str, prisma_client: PrismaClient | None):
1971 """
1972 - Check if user role in DB is admin
1973 - If not, update user role in DB to admin role
1974 """
1975 proxy_admin_id: Final = os.getenv("PROXY_ADMIN_ID")
1976 if proxy_admin_id is not None and proxy_admin_id == user_id:
1977 if user_role and user_role == LitellmUserRoles.PROXY_ADMIN.value:
1978 return user_role
1980 if prisma_client:
1981 await _user_meta_db(UserRepository(prisma_client)).update(
1982 where={"user_id": user_id},
1983 data={"user_role": LitellmUserRoles.PROXY_ADMIN.value},
1984 )
1986 user_role = LitellmUserRoles.PROXY_ADMIN.value
1988 return user_role
1991@router.get("/sso/callback", tags=["experimental"], include_in_schema=False)
1992async def auth_callback(request: Request, state: str | None = None):
1993 """Verify login"""
1994 verbose_proxy_logger.info("Starting SSO callback with state: %s", state)
1996 oauth_error: Final = request.query_params.get("error")
1997 if oauth_error:
1998 oauth_error_description: Final = request.query_params.get("error_description")
1999 verbose_proxy_logger.warning(
2000 "SSO callback received OAuth error: %s, description: %s", oauth_error, oauth_error_description
2001 )
2002 raise HTTPException(
2003 status_code=401,
2004 detail=f"OAuth error: {oauth_error}"
2005 + (f", error_description: {oauth_error_description}" if oauth_error_description else ""),
2006 )
2008 # Check if this is a CLI login (state starts with our CLI prefix)
2009 from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
2010 from litellm.proxy._types import LiteLLM_JWTAuth
2011 from litellm.proxy.auth.handle_jwt import JWTHandler
2012 from litellm.proxy.proxy_server import (
2013 general_settings,
2014 jwt_handler,
2015 master_key,
2016 prisma_client,
2017 user_api_key_cache,
2018 )
2020 if prisma_client is None:
2021 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
2023 sso_jwt_handler: JWTHandler | None = None
2024 ui_access_mode: Final = general_settings.get("ui_access_mode", None)
2025 if ui_access_mode is not None and isinstance(ui_access_mode, dict):
2026 sso_jwt_handler = JWTHandler()
2027 sso_jwt_handler.update_environment(
2028 prisma_client=prisma_client,
2029 user_api_key_cache=user_api_key_cache,
2030 litellm_jwtauth=LiteLLM_JWTAuth(
2031 team_ids_jwt_field=general_settings.get("ui_access_mode", {}).get("sso_group_jwt_field", None),
2032 ),
2033 leeway=0,
2034 )
2036 microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
2037 google_client_id: Final = os.getenv("GOOGLE_CLIENT_ID", None)
2038 generic_client_id: Final = os.getenv("GENERIC_CLIENT_ID", None)
2039 received_response: dict | None = None
2040 access_token_payload: dict | None = None
2041 sso_assertion: SSOIdentityAssertion | None = None
2042 # get url from request
2043 if master_key is None:
2044 raise ProxyException(
2045 message="Master Key not set for Proxy. Please set Master Key to use Admin UI. Set `LITELLM_MASTER_KEY` in .env or set general_settings:master_key in config.yaml. https://docs.litellm.ai/docs/proxy/virtual_keys. If set, use `--detailed_debug` to debug issue.",
2046 type=ProxyErrorTypes.auth_error,
2047 param="master_key",
2048 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
2049 )
2050 redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(request=request, sso_callback_route="sso/callback")
2052 verbose_proxy_logger.info("Redirecting to %s", redirect_url)
2053 result = None
2054 if google_client_id is not None:
2055 result = await GoogleSSOHandler.get_google_callback_response(
2056 request=request,
2057 google_client_id=google_client_id,
2058 redirect_url=redirect_url,
2059 )
2060 elif microsoft_client_id is not None:
2061 result = await MicrosoftSSOHandler.get_microsoft_callback_response(
2062 request=request,
2063 microsoft_client_id=microsoft_client_id,
2064 redirect_url=redirect_url,
2065 )
2067 elif generic_client_id is not None:
2068 (
2069 result,
2070 received_response,
2071 access_token_payload,
2072 sso_assertion,
2073 ) = await get_generic_sso_response(
2074 request=request,
2075 jwt_handler=jwt_handler,
2076 generic_client_id=generic_client_id,
2077 redirect_url=redirect_url,
2078 sso_jwt_handler=sso_jwt_handler,
2079 )
2081 if result is None:
2082 raise HTTPException(
2083 status_code=401,
2084 detail="Result not returned by SSO provider.",
2085 )
2087 if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
2088 # State format: {PREFIX}:{login_id}[:{user_code}]
2089 state_parts: Final = state.split(":", 2)
2090 key_id: Final = state_parts[1] if len(state_parts) > 1 else None
2091 prefill_user_code: Final = state_parts[2] if len(state_parts) > 2 else None
2093 verbose_proxy_logger.info("CLI SSO callback detected")
2094 return await cli_sso_callback(
2095 request=request,
2096 key=key_id,
2097 prefill_user_code=prefill_user_code,
2098 result=result,
2099 received_response=received_response,
2100 sso_assertion=sso_assertion,
2101 )
2103 # Control-plane cross-origin: read return_to from cookie.
2104 # Starlette's cookie_parser already handles RFC 2109 unquoting.
2105 cp_return_to: Final[str | None] = request.cookies.get("litellm_cp_return_to")
2107 return await SSOAuthenticationHandler.get_redirect_response_from_openid(
2108 result=result,
2109 request=request,
2110 received_response=received_response,
2111 generic_client_id=generic_client_id,
2112 ui_access_mode=ui_access_mode,
2113 access_token_payload=access_token_payload,
2114 jwt_handler=jwt_handler,
2115 return_to=cp_return_to,
2116 sso_assertion=sso_assertion,
2117 )
2120@router.get("/sso/saml/login", tags=["experimental"], include_in_schema=False)
2121async def saml_login(request: Request, return_to: str | None = None):
2122 """SP-initiated SAML login. Redirects the user to the configured IdP."""
2123 from litellm.proxy.proxy_server import user_api_key_cache
2125 _disable_ui_flag: Final = os.getenv("DISABLE_ADMIN_UI")
2126 if _disable_ui_flag is not None and str_to_bool(value=_disable_ui_flag):
2127 return admin_ui_disabled()
2129 return await SAMLAuthHandler.build_login_redirect(request=request, cache=user_api_key_cache, relay_state=return_to)
2132@router.get("/sso/saml/metadata", tags=["experimental"], include_in_schema=False)
2133async def saml_metadata(request: Request):
2134 """Service Provider metadata XML, for registering this proxy at the IdP."""
2135 from litellm.proxy.proxy_server import user_api_key_cache
2137 metadata: Final = await SAMLAuthHandler.build_sp_metadata(request=request, cache=user_api_key_cache)
2138 return Response(content=metadata, media_type="application/xml")
2141@router.post("/sso/saml/callback", tags=["experimental"], include_in_schema=False)
2142async def saml_callback(request: Request):
2143 """Assertion Consumer Service. Validates the IdP assertion and issues a UI session."""
2144 from litellm.proxy.proxy_server import (
2145 general_settings,
2146 jwt_handler,
2147 master_key,
2148 premium_user,
2149 prisma_client,
2150 user_api_key_cache,
2151 )
2153 _disable_ui_flag: Final = os.getenv("DISABLE_ADMIN_UI")
2154 if _disable_ui_flag is not None and str_to_bool(value=_disable_ui_flag):
2155 return admin_ui_disabled()
2157 if prisma_client is None:
2158 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
2159 if master_key is None:
2160 raise ProxyException(
2161 message="Master Key not set for Proxy. Set `LITELLM_MASTER_KEY` in .env or general_settings:master_key in config.yaml.",
2162 type=ProxyErrorTypes.auth_error,
2163 param="master_key",
2164 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
2165 )
2167 post_data: Final = await SAMLAuthHandler.read_acs_post_data(request)
2168 if "SAMLResponse" not in post_data:
2169 raise HTTPException(status_code=400, detail="Missing SAMLResponse in callback request.")
2171 result: Final = await SAMLAuthHandler.handle_acs(request=request, cache=user_api_key_cache, post_data=post_data)
2173 await _raise_if_sso_exceeds_free_user_limit(premium_user, prisma_client)
2175 ui_access_mode: Final = general_settings.get("ui_access_mode", None)
2176 relay_state: Final = post_data.get("RelayState")
2177 cp_return_to: Final[str | None] = (
2178 relay_state
2179 if isinstance(relay_state, str) and SSOAuthenticationHandler._validate_return_to(relay_state)
2180 else None
2181 )
2183 return await SSOAuthenticationHandler.get_redirect_response_from_openid(
2184 result=result,
2185 request=request,
2186 received_response=None,
2187 generic_client_id=None,
2188 ui_access_mode=ui_access_mode,
2189 access_token_payload=None,
2190 jwt_handler=jwt_handler,
2191 return_to=cp_return_to,
2192 )
2195async def _build_cli_sso_user_defined_values(
2196 result: OpenID | dict,
2197 parsed_openid_result: ParsedOpenIDResult,
2198) -> SSOUserDefinedValues | None:
2199 from litellm.proxy.proxy_server import user_custom_sso
2201 custom_sso_handler: Final[_CustomSsoCall | None] = user_custom_sso
2202 user_id: Final = parsed_openid_result.get("user_id")
2203 if custom_sso_handler is not None:
2204 if inspect.iscoroutinefunction(custom_sso_handler):
2205 return await custom_sso_handler(result)
2206 raise ValueError("user_custom_sso must be a coroutine function")
2207 if user_id is None:
2208 return None
2209 return SSOUserDefinedValues(
2210 models=[],
2211 user_id=user_id,
2212 user_email=parsed_openid_result.get("user_email"),
2213 max_budget=litellm.max_internal_user_budget,
2214 user_role=parsed_openid_result.get("user_role"),
2215 budget_duration=litellm.internal_user_budget_duration,
2216 )
2219def _cli_sso_team_detail(team_row: Mapping[str, object]) -> CliSsoTeamDetail:
2220 team: Final = _TeamRowGrants.model_validate(team_row)
2221 alias_table: Final = team.litellm_model_table
2222 return CliSsoTeamDetail(
2223 team_id=team.team_id,
2224 team_alias=team.team_alias,
2225 team_models=team.models,
2226 team_model_aliases=alias_table.model_aliases if alias_table is not None else None,
2227 )
2230async def fetch_cli_sso_team_details(
2231 prisma_client: PrismaClient,
2232 teams: Sequence[str],
2233) -> tuple[CliSsoTeamDetail, ...] | None:
2234 """``None`` means the lookup itself failed, which is not the same as the user having no teams."""
2235 if not teams:
2236 return ()
2237 try:
2238 prisma_teams: Final = await _team_detail_db(TeamRepository(prisma_client)).find_many(
2239 where={"team_id": {"in": teams}},
2240 include={"litellm_model_table": True},
2241 )
2242 except Exception as e:
2243 verbose_proxy_logger.error("Error fetching team details for CLI SSO session: %s", e)
2244 return None
2245 return tuple(_cli_sso_team_detail(team_row.model_dump()) for team_row in prisma_teams)
2248def _cli_sso_session_teams(team_details: Sequence[CliSsoTeamDetail]) -> list[str]:
2249 """The teams a login may bind to: only those whose row still exists.
2251 A team deleted out from under a membership, which is what deleting an organization
2252 leaves behind, can never resolve its grants, so offering it would refuse every
2253 future login for that user with nothing they could do to recover.
2254 """
2255 return [detail.team_id for detail in team_details if detail.team_id is not None]
2258def selected_cli_sso_team_detail(team_details: object, team_id: str | None) -> CliSsoTeamDetail | None:
2259 """``None`` means the team's grants are unknown. An empty grant is a real value meaning unrestricted,
2260 so an unknown one must not be minted as empty."""
2261 if team_id is None:
2262 return _TEAMLESS_CLI_SSO_TEAM_DETAIL
2263 try:
2264 details: Final = _CLI_SSO_TEAM_DETAILS_ADAPTER.validate_python(team_details)
2265 except ValidationError:
2266 return None
2267 return next((detail for detail in details if detail.team_id == team_id), None)
2270async def _complete_cli_sso_callback_session(
2271 *,
2272 request: Request,
2273 key: str,
2274 flow: dict,
2275 result: OpenID | dict,
2276 parsed_openid_result: ParsedOpenIDResult,
2277 user_defined_values: SSOUserDefinedValues | None,
2278 prisma_client: PrismaClient,
2279 user_api_key_cache: UserApiKeyCache,
2280 cli_sso_session_cache: DualCache,
2281 proxy_logging_obj: ProxyLogging,
2282 prefill_user_code: str | None = None,
2283 sso_assertion: SSOIdentityAssertion | None = None,
2284):
2285 from fastapi.responses import HTMLResponse
2287 user_id: Final = parsed_openid_result.get("user_id")
2288 user_email: Final = parsed_openid_result.get("user_email")
2289 user_info: Final = await get_user_info_from_db(
2290 result=result,
2291 prisma_client=prisma_client,
2292 user_api_key_cache=user_api_key_cache,
2293 proxy_logging_obj=proxy_logging_obj,
2294 user_email=user_email,
2295 user_defined_values=user_defined_values,
2296 alternate_user_id=user_id,
2297 )
2298 if user_info is None:
2299 raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO")
2300 if not user_info.user_id:
2301 raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO")
2303 await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion)
2304 await warn_if_id_jag_assertion_uncaptured(sso_assertion)
2306 teams: list[str] = []
2307 if hasattr(user_info, "teams") and user_info.teams:
2308 teams = user_info.teams if isinstance(user_info.teams, list) else []
2310 team_details: Final = await fetch_cli_sso_team_details(prisma_client=prisma_client, teams=teams)
2311 if team_details is None:
2312 raise HTTPException(
2313 status_code=500,
2314 detail="Could not resolve team model grants for this login. Please try again",
2315 )
2316 resolved_teams: Final = _cli_sso_session_teams(team_details)
2317 attribution_metadata: Final = build_cli_sso_attribution_metadata(result=result)
2318 if attribution_metadata:
2319 await _persist_cli_sso_user_metadata(
2320 prisma_client=prisma_client,
2321 user_id=cast(str, user_info.user_id),
2322 attribution_metadata=attribution_metadata,
2323 )
2325 flow["session_data"] = {
2326 "user_id": cast(str, user_info.user_id),
2327 "user_role": user_info.user_role,
2328 "models": user_info.models if hasattr(user_info, "models") else [],
2329 "user_email": user_email,
2330 "teams": resolved_teams,
2331 "team_details": [detail.model_dump() for detail in team_details],
2332 "attribution_metadata": attribution_metadata,
2333 }
2334 flow["sso_complete"] = True
2335 browser_complete_token: Final = secrets.token_urlsafe(32)
2336 flow["browser_complete_token_hash"] = _hash_cli_sso_secret(browser_complete_token)
2337 _set_cli_sso_flow(login_id=key, cache=cli_sso_session_cache, flow=flow)
2339 verbose_proxy_logger.info(
2340 "Stored CLI SSO session for user: %s, teams: %s, num_teams: %s",
2341 user_info.user_id,
2342 resolved_teams,
2343 len(resolved_teams),
2344 )
2345 verify_url: Final = get_custom_url(
2346 request_base_url=str(request.base_url),
2347 route=f"sso/cli/complete/{key}",
2348 )
2349 return HTMLResponse(
2350 content=_render_cli_sso_verification_page(
2351 verify_url=verify_url,
2352 browser_complete_token=browser_complete_token,
2353 prefill_user_code=prefill_user_code,
2354 ),
2355 status_code=200,
2356 )
2359async def cli_sso_callback(
2360 request: Request,
2361 key: str | None = None,
2362 result: OpenID | dict | None = None,
2363 received_response: dict | None = None,
2364 prefill_user_code: str | None = None,
2365 sso_assertion: SSOIdentityAssertion | None = None,
2366):
2367 """CLI SSO callback - stores session info for JWT generation on polling"""
2368 verbose_proxy_logger.info("CLI SSO callback")
2370 from litellm.proxy.proxy_server import (
2371 cli_sso_session_cache,
2372 general_settings,
2373 prisma_client,
2374 proxy_logging_obj,
2375 user_api_key_cache,
2376 )
2378 flow: Final = _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache)
2380 if prisma_client is None:
2381 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
2383 if result is None:
2384 raise HTTPException(
2385 status_code=500,
2386 detail="SSO authentication failed - no result returned from provider",
2387 )
2389 # After None check, cast to non-None type for type checker
2390 result_non_none: Final[OpenID | dict] = cast(OpenID | dict, result)
2392 try:
2393 parsed_openid_result: Final = SSOAuthenticationHandler._get_user_email_and_id_from_result(
2394 result=result_non_none,
2395 generic_client_id=os.getenv("GENERIC_CLIENT_ID", None),
2396 )
2397 verbose_proxy_logger.debug("parsed_openid_result: %s", parsed_openid_result)
2398 user_defined_values: Final = await _build_cli_sso_user_defined_values(
2399 result=result_non_none,
2400 parsed_openid_result=parsed_openid_result,
2401 )
2403 SSOAuthenticationHandler.verify_user_in_restricted_sso_group(
2404 general_settings=general_settings,
2405 result=result_non_none,
2406 received_response=received_response,
2407 )
2409 return await _complete_cli_sso_callback_session(
2410 request=request,
2411 key=cast(str, key),
2412 flow=flow,
2413 result=result_non_none,
2414 parsed_openid_result=parsed_openid_result,
2415 user_defined_values=user_defined_values,
2416 prisma_client=prisma_client,
2417 user_api_key_cache=user_api_key_cache,
2418 cli_sso_session_cache=cli_sso_session_cache,
2419 proxy_logging_obj=proxy_logging_obj,
2420 prefill_user_code=prefill_user_code,
2421 sso_assertion=sso_assertion,
2422 )
2423 except ProxyException:
2424 raise
2425 except HTTPException:
2426 raise
2427 except Exception as e:
2428 verbose_proxy_logger.error("Error with CLI SSO callback: %s", e)
2429 raise HTTPException(status_code=500, detail=f"Failed to process CLI SSO: {e}")
2432@router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False)
2433async def cli_poll_key(
2434 key_id: str,
2435 team_id: str | None = None,
2436 x_litellm_cli_poll_secret: str | None = Header(default=None),
2437):
2438 """
2439 CLI polling endpoint - retrieves session from cache and generates JWT.
2441 Flow:
2442 1. First poll (no team_id): Returns teams list without generating JWT
2443 2. Second poll (with team_id): Generates JWT with selected team and deletes session
2445 Args:
2446 key_id: The CLI login session ID
2447 team_id: Optional team ID to assign to the JWT. If provided, must be one of user's teams.
2448 """
2449 from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
2450 from litellm.proxy.proxy_server import cli_sso_session_cache
2452 try:
2453 flow: Final = _get_cli_sso_flow_or_raise(login_id=key_id, cache=cli_sso_session_cache)
2454 if not _verify_cli_sso_poll_secret(flow=flow, poll_secret=x_litellm_cli_poll_secret):
2455 raise HTTPException(status_code=403, detail="Invalid CLI polling secret")
2457 if not flow.get("sso_complete") or not flow.get("user_code_verified"):
2458 return {"status": "pending"}
2460 session_data: Final = flow.get("session_data")
2462 if isinstance(session_data, dict):
2463 user_teams: Final = session_data.get("teams", [])
2464 user_team_details: Final = session_data.get("team_details")
2465 user_id: Final = session_data["user_id"]
2467 verbose_proxy_logger.info(
2468 "CLI poll: user=%s, team_id=%s, user_teams=%s, num_teams=%s",
2469 user_id,
2470 team_id,
2471 user_teams,
2472 len(user_teams),
2473 )
2475 # If no team_id provided and user has teams, return teams list for selection
2476 # Don't generate JWT yet - let CLI select a team first. For newer
2477 # clients we return rich team details (id + alias); older clients
2478 # can continue to rely on the simple "teams" list.
2479 if team_id is None and len(user_teams) > 1:
2480 verbose_proxy_logger.info("Returning teams list for user %s to select from: %s", user_id, user_teams)
2481 # Best-effort construction of team_details if it wasn't
2482 # already cached for some reason.
2483 team_details_response: list[dict[str, object]] | None = None
2484 if isinstance(user_team_details, list) and user_team_details:
2485 team_details_response = user_team_details
2486 elif user_teams:
2487 team_details_response = [{"team_id": t, "team_alias": None} for t in user_teams]
2488 poll_response: dict[str, object] = {
2489 "status": "ready",
2490 "user_id": user_id,
2491 "teams": user_teams,
2492 "team_details": team_details_response,
2493 "requires_team_selection": True,
2494 }
2495 attribution_metadata = _cli_poll_attribution_metadata_from_session(session_data)
2496 if attribution_metadata:
2497 poll_response["attribution_metadata"] = attribution_metadata
2498 return poll_response
2500 # Validate team_id if provided
2501 if team_id is not None:
2502 if team_id not in user_teams:
2503 raise HTTPException(
2504 status_code=403,
2505 detail=f"User does not belong to team: {team_id}. Available teams: {user_teams}",
2506 )
2507 else:
2508 # If no team_id provided and user has 0 or 1 team, use first team (or None)
2509 team_id = user_teams[0] if len(user_teams) > 0 else None
2511 selected_team: Final = selected_cli_sso_team_detail(
2512 team_details=user_team_details,
2513 team_id=team_id,
2514 )
2515 if selected_team is None:
2516 raise HTTPException(
2517 status_code=500,
2518 detail=f"Could not resolve the model grants for team: {team_id}. Please run `lite login` again",
2519 )
2521 user_info: Final = LiteLLM_UserTable(
2522 user_id=user_id,
2523 user_role=session_data["user_role"],
2524 models=session_data.get("models", []),
2525 )
2527 jwt_token: Final = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
2528 user_info=user_info,
2529 team_id=team_id,
2530 team_alias=selected_team.team_alias,
2531 team_models=selected_team.team_models,
2532 team_model_aliases=selected_team.team_model_aliases,
2533 max_budget=None,
2534 )
2536 # Delete cache entry (single-use)
2537 cli_sso_session_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id))
2539 verbose_proxy_logger.info("CLI JWT generated for user: %s, team: %s", user_id, team_id)
2540 poll_response = {
2541 "status": "ready",
2542 "key": jwt_token,
2543 "user_id": user_id,
2544 "team_id": team_id,
2545 "teams": user_teams,
2546 # Echo back any team details we have so clients can
2547 # present nicer information if needed.
2548 "team_details": user_team_details,
2549 }
2550 attribution_metadata = _cli_poll_attribution_metadata_from_session(session_data)
2551 if attribution_metadata:
2552 poll_response["attribution_metadata"] = attribution_metadata
2553 return poll_response
2554 else:
2555 return {"status": "pending"}
2557 except HTTPException:
2558 raise
2559 except Exception as e:
2560 verbose_proxy_logger.error("Error polling for CLI JWT: %s", e)
2561 raise HTTPException(status_code=500, detail=f"Error checking session status: {e}")
2564async def insert_sso_user(
2565 result_openid: OpenID | dict | None,
2566 user_defined_values: SSOUserDefinedValues | None = None,
2567) -> NewUserResponse:
2568 """
2569 Helper function to create a New User in LiteLLM DB after a successful SSO login
2571 Args:
2572 result_openid (OpenID): User information in OpenID format if the login was successful.
2573 user_defined_values (Optional[SSOUserDefinedValues], optional): LiteLLM SSOValues / fields that were read
2575 Returns:
2576 Tuple[str, str]: User ID and User Role
2577 """
2578 verbose_proxy_logger.debug("Inserting SSO user into DB. User values: %s", user_defined_values)
2579 if result_openid is None:
2580 raise ValueError("result_openid is None")
2581 if isinstance(result_openid, dict):
2582 result_openid = OpenID(**result_openid)
2584 if user_defined_values is None:
2585 raise ValueError("user_defined_values is None")
2587 # Apply default_internal_user_params
2588 if litellm.default_internal_user_params:
2589 # Preserve the SSO-extracted role if it's a valid LiteLLM role,
2590 # regardless of how it was determined (role_mappings, Microsoft app_roles,
2591 # GENERIC_USER_ROLE_ATTRIBUTE, custom SSO handler, etc.)
2592 sso_role: Final = user_defined_values.get("user_role")
2593 if _should_use_role_from_sso_response(sso_role):
2594 # Preserve the SSO-extracted role, but apply other defaults
2595 preserved_role: Final = sso_role
2596 user_defined_values.update(litellm.default_internal_user_params)
2597 user_defined_values["user_role"] = preserved_role # Restore preserved role
2598 verbose_proxy_logger.debug("Preserved SSO-extracted role '%s'", preserved_role)
2599 else:
2600 # SSO didn't provide a valid role, apply all defaults including role
2601 user_defined_values.update(litellm.default_internal_user_params)
2603 # Set budget for internal users
2604 if user_defined_values.get("user_role") == LitellmUserRoles.INTERNAL_USER.value:
2605 if user_defined_values.get("max_budget") is None:
2606 user_defined_values["max_budget"] = litellm.max_internal_user_budget
2607 if user_defined_values.get("budget_duration") is None:
2608 user_defined_values["budget_duration"] = litellm.internal_user_budget_duration
2610 if user_defined_values["user_role"] is None:
2611 user_defined_values["user_role"] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
2613 new_user_request: Final = NewUserRequest(
2614 user_id=user_defined_values["user_id"],
2615 user_email=normalize_email(user_defined_values["user_email"]),
2616 user_role=user_defined_values["user_role"],
2617 max_budget=user_defined_values["max_budget"],
2618 budget_duration=user_defined_values["budget_duration"],
2619 sso_user_id=user_defined_values["user_id"],
2620 auto_create_key=False,
2621 )
2623 if result_openid and hasattr(result_openid, "provider"):
2624 new_user_request.metadata = {"auth_provider": getattr(result_openid, "provider")}
2626 response: Final = await new_user(
2627 data=new_user_request,
2628 user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
2629 )
2631 return response
2634@router.get(
2635 "/sso/get/ui_settings",
2636 tags=["experimental"],
2637 include_in_schema=False,
2638 dependencies=[Depends(user_api_key_auth)],
2639)
2640async def get_ui_settings(request: Request):
2641 from litellm.proxy.proxy_server import general_settings, proxy_state
2643 _proxy_base_url: Final = os.getenv("PROXY_BASE_URL", None)
2644 _logout_url: Final = os.getenv("PROXY_LOGOUT_URL", None)
2645 _api_doc_base_url: Final = os.getenv("LITELLM_UI_API_DOC_BASE_URL", None)
2646 _is_sso_enabled: Final = has_user_setup_sso()
2647 disable_expensive_db_queries: Final = (
2648 proxy_state.get_proxy_state_variable("spend_logs_row_count") > MAX_SPENDLOG_ROWS_TO_QUERY
2649 )
2650 default_team_disabled = general_settings.get("default_team_disabled", False)
2651 if "PROXY_DEFAULT_TEAM_DISABLED" in os.environ:
2652 if os.environ["PROXY_DEFAULT_TEAM_DISABLED"].lower() == "true":
2653 default_team_disabled = True
2655 return {
2656 "PROXY_BASE_URL": _proxy_base_url,
2657 "PROXY_LOGOUT_URL": _logout_url,
2658 "LITELLM_UI_API_DOC_BASE_URL": _api_doc_base_url,
2659 "DEFAULT_TEAM_DISABLED": default_team_disabled,
2660 "SSO_ENABLED": _is_sso_enabled,
2661 "NUM_SPEND_LOGS_ROWS": proxy_state.get_proxy_state_variable("spend_logs_row_count"),
2662 "DISABLE_EXPENSIVE_DB_QUERIES": disable_expensive_db_queries,
2663 }
2666@router.get(
2667 "/sso/readiness",
2668 tags=["experimental"],
2669 dependencies=[Depends(user_api_key_auth)],
2670)
2671async def sso_readiness():
2672 """
2673 Health endpoint for checking SSO readiness.
2674 Checks if the configured SSO provider has all required environment variables set in memory.
2675 """
2676 microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
2677 google_client_id: Final = os.getenv("GOOGLE_CLIENT_ID", None)
2678 generic_client_id: Final = os.getenv("GENERIC_CLIENT_ID", None)
2680 # Determine which SSO provider is configured
2681 configured_provider = None
2682 if google_client_id is not None: 2682 ↛ 2683line 2682 didn't jump to line 2683 because the condition on line 2682 was never true
2683 configured_provider = "google"
2684 elif microsoft_client_id is not None: 2684 ↛ 2685line 2684 didn't jump to line 2685 because the condition on line 2684 was never true
2685 configured_provider = "microsoft"
2686 elif generic_client_id is not None: 2686 ↛ 2687line 2686 didn't jump to line 2687 because the condition on line 2686 was never true
2687 configured_provider = "generic"
2689 # If no SSO is configured, return healthy (SSO is optional)
2690 if configured_provider is None: 2690 ↛ 2698line 2690 didn't jump to line 2698 because the condition on line 2690 was always true
2691 return {
2692 "status": "healthy",
2693 "sso_configured": False,
2694 "message": "No SSO provider configured",
2695 }
2697 # Check required environment variables for the configured provider
2698 missing_vars: Final = []
2700 if configured_provider == "google":
2701 google_client_secret: Final = os.getenv("GOOGLE_CLIENT_SECRET", None)
2702 if google_client_secret is None:
2703 missing_vars.append("GOOGLE_CLIENT_SECRET")
2705 elif configured_provider == "microsoft":
2706 microsoft_client_secret: Final = os.getenv("MICROSOFT_CLIENT_SECRET", None)
2707 microsoft_tenant: Final = os.getenv("MICROSOFT_TENANT", None)
2708 if microsoft_client_secret is None:
2709 missing_vars.append("MICROSOFT_CLIENT_SECRET")
2710 if microsoft_tenant is None:
2711 missing_vars.append("MICROSOFT_TENANT")
2713 elif configured_provider == "generic":
2714 generic_client_secret: Final = os.getenv("GENERIC_CLIENT_SECRET", None)
2715 generic_authorization_endpoint: Final = os.getenv("GENERIC_AUTHORIZATION_ENDPOINT", None)
2716 generic_token_endpoint: Final = os.getenv("GENERIC_TOKEN_ENDPOINT", None)
2717 generic_userinfo_endpoint: Final = os.getenv("GENERIC_USERINFO_ENDPOINT", None)
2718 if generic_client_secret is None:
2719 missing_vars.append("GENERIC_CLIENT_SECRET")
2720 if generic_authorization_endpoint is None:
2721 missing_vars.append("GENERIC_AUTHORIZATION_ENDPOINT")
2722 if generic_token_endpoint is None:
2723 missing_vars.append("GENERIC_TOKEN_ENDPOINT")
2724 if generic_userinfo_endpoint is None:
2725 missing_vars.append("GENERIC_USERINFO_ENDPOINT")
2727 # If all required variables are present, return healthy
2728 if len(missing_vars) == 0:
2729 return {
2730 "status": "healthy",
2731 "sso_configured": True,
2732 "provider": configured_provider,
2733 "message": f"{configured_provider.capitalize()} SSO is properly configured",
2734 }
2736 # If some variables are missing, return unhealthy
2737 raise HTTPException(
2738 status_code=503,
2739 detail={
2740 "status": "unhealthy",
2741 "sso_configured": True,
2742 "provider": configured_provider,
2743 "missing_environment_variables": missing_vars,
2744 "message": f"{configured_provider.capitalize()} SSO is configured but missing required environment variables: {', '.join(missing_vars)}",
2745 },
2746 )
2749def _is_same_origin_return_path(return_to: str) -> bool:
2750 """True for a strictly relative return path that stays on the gateway's own origin by
2751 construction, and is therefore safe to honor without a configured ``control_plane_url``.
2752 Used by the MCP gateway DCR authorize round-trip so a browser sent through login lands
2753 back on the authorize request.
2755 Requires a single leading ``/`` (not protocol-relative ``//``), no backslash (browsers
2756 fold ``\\`` to ``/``, so ``/\\evil.com`` would escape the origin), and no control or
2757 whitespace characters. Rejecting control chars keeps a ``\\r\\n``/tab-bearing value out
2758 of the redirect ``Location`` and the ``litellm_cp_return_to`` cookie entirely, rather
2759 than relying on downstream header encoding to neutralize it."""
2760 if not return_to.startswith("/") or return_to.startswith("//") or "\\" in return_to:
2761 return False
2762 return not any(ord(ch) < 0x20 or ch in (" ", "\x7f") for ch in return_to)
2765async def _sso_return_to_redirect(
2766 return_to: str | None,
2767 jwt_token: str,
2768 redis_usage_cache,
2769 user_api_key_cache,
2770 request: Request,
2771) -> RedirectResponse | None:
2772 """Resolve the post-SSO redirect for a ``return_to``, or None to fall through to the dashboard.
2774 Two arms, both clearing the one-shot ``litellm_cp_return_to`` cookie:
2775 - **Same-origin relative path** (the MCP gateway DCR authorize round-trip): set the session cookie
2776 exactly like the dashboard path, then send the browser back where it came from.
2777 - **Control-plane cross-origin** (``control_plane_url``): stash the JWT behind a single-use opaque
2778 code (60s TTL) so the token never lands in browser history/logs; the control plane redeems it via
2779 ``POST /v3/login/exchange``.
2781 Extracted from ``get_redirect_response_from_openid`` to keep that method inside the complexity
2782 budget; behavior is identical to the inline arms it replaces (including letting
2783 ``_validate_return_to`` raise for a mismatched absolute return_to, as before)."""
2784 if return_to is None:
2785 return None
2787 if _is_same_origin_return_path(return_to):
2788 redirect_response = RedirectResponse(url=return_to, status_code=303)
2789 set_session_token_cookie(redirect_response, request, jwt_token)
2790 redirect_response.delete_cookie("litellm_cp_return_to")
2791 return redirect_response
2793 if SSOAuthenticationHandler._validate_return_to(return_to):
2794 code: Final = secrets.token_urlsafe(32)
2795 cache_key: Final = f"login_code:{code}"
2796 cache_value: Final = {"token": jwt_token, "redirect_url": return_to}
2797 if redis_usage_cache is not None:
2798 await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
2799 else:
2800 await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
2802 separator: Final = "&" if "?" in return_to else "?"
2803 redirect_url: Final = return_to + separator + urlencode({"login": "success", "code": code})
2804 verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code")
2805 redirect_response = RedirectResponse(url=redirect_url, status_code=303)
2806 redirect_response.delete_cookie("litellm_cp_return_to")
2807 return redirect_response
2809 return None
2812def set_session_token_cookie(response: Response, request: Request, jwt_token: str) -> None:
2813 """Set the ``token`` session cookie shared by every sign-in path.
2815 Not HttpOnly: the dashboard reads this cookie via ``document.cookie`` to
2816 populate its own Authorization headers (see
2817 ``ui/litellm-dashboard/src/utils/cookieUtils.ts``), so marking it
2818 HttpOnly would break login. Secure is still required whenever the public
2819 origin is HTTPS, resolved the same trust-aware way as every other
2820 litellm cookie."""
2821 response.set_cookie(
2822 key="token",
2823 value=jwt_token,
2824 secure=IPAddressUtils.is_request_https(request),
2825 httponly=False,
2826 samesite="lax",
2827 )
2830def _persist_return_to_cookie(response: Response, return_to: str | None, request: Request) -> None:
2831 """Best-effort: persist a SAFE ``return_to`` on ``response`` as the one-shot ``litellm_cp_return_to``
2832 cookie so ANY sign-in path — SSO / Okta / generic OR the username/password form — can resume there
2833 afterwards. THIS is the single source of truth, called by every sign-in branch so they cannot
2834 diverge (a per-branch reimplementation is exactly how the two drifted before). Honors a strictly
2835 relative same-origin path, and (when ``control_plane_url`` is configured) a return_to matching that
2836 origin. It NEVER raises: a mismatched or invalid ``return_to`` is simply not stored, so it can never
2837 block sign-in — the login entrypoint must always render."""
2838 if return_to is None:
2839 return
2840 try:
2841 safe: Final = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler._validate_return_to(return_to)
2842 except HTTPException:
2843 return # a non-matching absolute return_to is ignored, never blocks sign-in
2844 if safe:
2845 response.set_cookie(
2846 key="litellm_cp_return_to",
2847 value=return_to,
2848 max_age=600,
2849 httponly=True,
2850 samesite="lax",
2851 secure=IPAddressUtils.is_request_https(request),
2852 )
2855class SSOAuthenticationHandler:
2856 """
2857 Handler for SSO Authentication across all SSO providers
2858 """
2860 @staticmethod
2861 def _validate_return_to(return_to: str) -> bool:
2862 """
2863 Validate that return_to matches the configured control_plane_url origin.
2865 Returns True if return_to is valid and should be used.
2866 Returns False if control_plane_url is not configured (return_to is ignored).
2867 Raises HTTPException(400) if return_to origin does not match control_plane_url origin.
2868 """
2869 from litellm.proxy.proxy_server import general_settings
2871 control_plane_url: Final = general_settings.get("control_plane_url")
2872 if control_plane_url is None:
2873 return False
2875 def _origin(url: str) -> tuple:
2876 parsed: Final = urlparse(url)
2877 scheme: Final = (parsed.scheme or "").lower()
2878 hostname: Final = (parsed.hostname or "").lower()
2879 default_port: Final = 443 if scheme == "https" else 80
2880 port: Final = parsed.port if parsed.port is not None else default_port
2881 return (scheme, hostname, port)
2883 if _origin(return_to) != _origin(control_plane_url):
2884 raise HTTPException(
2885 status_code=400,
2886 detail="return_to does not match the configured control_plane_url",
2887 )
2889 return True
2891 @staticmethod
2892 async def get_sso_login_redirect(
2893 redirect_url: str,
2894 google_client_id: str | None = None,
2895 microsoft_client_id: str | None = None,
2896 generic_client_id: str | None = None,
2897 state: str | None = None,
2898 request: Request | None = None,
2899 ) -> RedirectResponse | None:
2900 """
2901 Step 1. Call Get Login Redirect for the SSO provider. Send the redirect response to `redirect_url`
2903 Args:
2904 redirect_url (str): The URL to redirect the user to after login
2905 google_client_id (Optional[str], optional): The Google Client ID. Defaults to None.
2906 microsoft_client_id (Optional[str], optional): The Microsoft Client ID. Defaults to None.
2907 generic_client_id (Optional[str], optional): The Generic Client ID. Defaults to None.
2908 request: Optional FastAPI request, used to drive the ``Secure``
2909 attribute on the ``litellm_oauth_state`` CSRF cookie.
2911 Returns:
2912 RedirectResponse: The redirect response from the SSO provider.
2913 """
2914 # Google SSO Auth
2915 if google_client_id is not None:
2916 from fastapi_sso.sso.google import GoogleSSO
2918 google_client_secret: Final = os.getenv("GOOGLE_CLIENT_SECRET", None)
2919 if google_client_secret is None:
2920 raise ProxyException(
2921 message="GOOGLE_CLIENT_SECRET not set. Set it in .env file",
2922 type=ProxyErrorTypes.auth_error,
2923 param="GOOGLE_CLIENT_SECRET",
2924 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
2925 )
2926 google_sso: Final = GoogleSSO(
2927 client_id=google_client_id,
2928 client_secret=google_client_secret,
2929 redirect_uri=redirect_url,
2930 )
2931 verbose_proxy_logger.info(
2932 "In /google-login/key/generate, \nGOOGLE_REDIRECT_URI: %s\nGOOGLE_CLIENT_ID: %s",
2933 redirect_url,
2934 google_client_id,
2935 )
2936 with google_sso:
2937 return await google_sso.get_login_redirect(state=state)
2938 # Microsoft SSO Auth
2939 elif microsoft_client_id is not None:
2940 microsoft_client_secret: Final = os.getenv("MICROSOFT_CLIENT_SECRET", None)
2941 microsoft_tenant: Final = os.getenv("MICROSOFT_TENANT", None)
2942 if microsoft_client_secret is None:
2943 raise ProxyException(
2944 message="MICROSOFT_CLIENT_SECRET not set. Set it in .env file",
2945 type=ProxyErrorTypes.auth_error,
2946 param="MICROSOFT_CLIENT_SECRET",
2947 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
2948 )
2949 microsoft_sso: Final = CustomMicrosoftSSO(
2950 client_id=microsoft_client_id,
2951 client_secret=microsoft_client_secret,
2952 tenant=microsoft_tenant,
2953 redirect_uri=redirect_url,
2954 allow_insecure_http=True,
2955 )
2956 with microsoft_sso:
2957 return await microsoft_sso.get_login_redirect(state=state)
2958 elif generic_client_id is not None:
2959 from fastapi_sso.sso.base import DiscoveryDocument
2960 from fastapi_sso.sso.generic import create_provider
2962 generic_client_secret: Final = os.getenv("GENERIC_CLIENT_SECRET", None)
2963 generic_scope: Final = os.getenv("GENERIC_SCOPE", "openid email profile").split(" ")
2964 generic_authorization_endpoint: Final = os.getenv("GENERIC_AUTHORIZATION_ENDPOINT", None)
2965 generic_token_endpoint: Final = os.getenv("GENERIC_TOKEN_ENDPOINT", None)
2966 generic_userinfo_endpoint: Final = os.getenv("GENERIC_USERINFO_ENDPOINT", None)
2967 if generic_client_secret is None:
2968 raise ProxyException(
2969 message="GENERIC_CLIENT_SECRET not set. Set it in .env file",
2970 type=ProxyErrorTypes.auth_error,
2971 param="GENERIC_CLIENT_SECRET",
2972 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
2973 )
2974 if generic_authorization_endpoint is None:
2975 raise ProxyException(
2976 message="GENERIC_AUTHORIZATION_ENDPOINT not set. Set it in .env file",
2977 type=ProxyErrorTypes.auth_error,
2978 param="GENERIC_AUTHORIZATION_ENDPOINT",
2979 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
2980 )
2981 if generic_token_endpoint is None:
2982 raise ProxyException(
2983 message="GENERIC_TOKEN_ENDPOINT not set. Set it in .env file",
2984 type=ProxyErrorTypes.auth_error,
2985 param="GENERIC_TOKEN_ENDPOINT",
2986 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
2987 )
2988 if generic_userinfo_endpoint is None:
2989 raise ProxyException(
2990 message="GENERIC_USERINFO_ENDPOINT not set. Set it in .env file",
2991 type=ProxyErrorTypes.auth_error,
2992 param="GENERIC_USERINFO_ENDPOINT",
2993 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
2994 )
2995 verbose_proxy_logger.debug(
2996 "authorization_endpoint: %s\ntoken_endpoint: %s\nuserinfo_endpoint: %s",
2997 generic_authorization_endpoint,
2998 generic_token_endpoint,
2999 generic_userinfo_endpoint,
3000 )
3001 verbose_proxy_logger.debug(
3002 "GENERIC_REDIRECT_URI: %s\nGENERIC_CLIENT_ID: %s\n", redirect_url, generic_client_id
3003 )
3004 discovery: Final = DiscoveryDocument(
3005 authorization_endpoint=generic_authorization_endpoint,
3006 token_endpoint=generic_token_endpoint,
3007 userinfo_endpoint=generic_userinfo_endpoint,
3008 )
3009 SSOProvider: Final = create_provider(name="oidc", discovery_document=discovery)
3010 generic_sso: Final = SSOProvider(
3011 client_id=generic_client_id,
3012 client_secret=generic_client_secret,
3013 redirect_uri=redirect_url,
3014 allow_insecure_http=True,
3015 scope=generic_scope,
3016 )
3017 return await SSOAuthenticationHandler.get_generic_sso_redirect_response(
3018 generic_sso=generic_sso,
3019 state=state,
3020 generic_authorization_endpoint=generic_authorization_endpoint,
3021 request=request,
3022 )
3023 raise ValueError(
3024 "Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso"
3025 )
3027 @staticmethod
3028 async def get_generic_sso_redirect_response(
3029 generic_sso: Any,
3030 state: str | None = None,
3031 generic_authorization_endpoint: str | None = None,
3032 request: Request | None = None,
3033 ) -> RedirectResponse | None:
3034 """
3035 Get the redirect response for Generic SSO
3036 """
3037 from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
3039 from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
3041 with generic_sso:
3042 # State is bound to the caller's browser via a ``litellm_oauth_state``
3043 # HttpOnly cookie set on the redirect response below; the SSO
3044 # callback validates the URL ``state`` against that cookie before
3045 # completing the PKCE token exchange. Without this binding, an
3046 # attacker who pre-mints a state + a cached PKCE verifier can hand
3047 # the link to a victim and capture the resulting access token
3048 # (Login CSRF / token theft).
3049 (
3050 redirect_params,
3051 code_verifier,
3052 ) = SSOAuthenticationHandler._get_generic_sso_redirect_params(
3053 state=state,
3054 generic_authorization_endpoint=generic_authorization_endpoint,
3055 )
3057 # Separate PKCE params from state params (fastapi-sso doesn't accept code_challenge)
3058 pkce_params: Final = {}
3059 state_only_params: Final = {}
3060 for key, value in redirect_params.items():
3061 if key in ("code_challenge", "code_challenge_method"):
3062 pkce_params[key] = value
3063 else:
3064 state_only_params[key] = value
3066 # Get the redirect response from fastapi-sso with only state param
3067 redirect_response: Final = await generic_sso.get_login_redirect(**state_only_params)
3069 # If PKCE is enabled, add PKCE parameters to the redirect URL
3070 if code_verifier and "state" in redirect_params:
3071 # Store code_verifier in cache (10 min TTL). Wrap in dict for proper
3072 # JSON serialization in Redis. Use Redis when available so callbacks
3073 # landing on another pod can retrieve it (multi-pod SSO).
3074 cache_key: Final = f"pkce_verifier:{redirect_params['state']}"
3075 if redis_usage_cache is not None:
3076 await redis_usage_cache.async_set_cache(
3077 key=cache_key,
3078 value={"code_verifier": code_verifier},
3079 ttl=600,
3080 )
3081 else:
3082 await user_api_key_cache.async_set_cache(
3083 key=cache_key,
3084 value={"code_verifier": code_verifier},
3085 ttl=600,
3086 )
3087 verbose_proxy_logger.debug("PKCE code_verifier stored in cache (TTL: 600s)")
3089 # Add PKCE parameters to the authorization URL
3090 if pkce_params:
3091 parsed_url: Final = urlparse(str(redirect_response.headers["location"]))
3092 query_params: Final = parse_qs(parsed_url.query)
3094 # Add PKCE parameters
3095 for key, value in pkce_params.items():
3096 query_params[key] = [value]
3098 # Reconstruct the URL with PKCE parameters
3099 new_query: Final = urlencode(query_params, doseq=True)
3100 new_url: Final = urlunparse(
3101 (
3102 parsed_url.scheme,
3103 parsed_url.netloc,
3104 parsed_url.path,
3105 parsed_url.params,
3106 new_query,
3107 parsed_url.fragment,
3108 )
3109 )
3111 # Update the redirect response
3112 redirect_response.headers["location"] = new_url
3114 # Bind state to the user's browser session. The /callback
3115 # handler validates the URL ``state`` against this cookie via
3116 # ``secrets.compare_digest`` before exchanging the PKCE
3117 # code_verifier. Only set the cookie when PKCE is in use
3118 # (i.e. inside this ``code_verifier`` branch) so two
3119 # concurrent SSO sessions — one PKCE, one plain — cannot
3120 # overwrite each other's state cookie.
3121 state_value: Final = redirect_params.get("state")
3122 if state_value and redirect_response is not None:
3123 # Production-safe default: require HTTPS for the
3124 # CSRF-protection cookie unless we can prove the
3125 # incoming request is HTTP (local dev). Without
3126 # ``Secure`` the cookie is sent over plain HTTP,
3127 # letting a network observer read and replay the
3128 # state value and bypass this protection. Trust-aware:
3129 # honors PROXY_BASE_URL / a trusted reverse proxy's
3130 # X-Forwarded-Proto instead of only the literal scheme
3131 # litellm sees on the wire.
3132 secure_flag: Final = request is None or IPAddressUtils.is_request_https(request)
3133 redirect_response.set_cookie(
3134 key="litellm_oauth_state",
3135 value=state_value,
3136 max_age=600,
3137 httponly=True,
3138 samesite="lax",
3139 secure=secure_flag,
3140 )
3141 return redirect_response
3143 @staticmethod
3144 def _get_generic_sso_redirect_params(
3145 state: str | None = None,
3146 generic_authorization_endpoint: str | None = None,
3147 ) -> tuple[dict[str, str], str | None]:
3148 """
3149 Get redirect parameters for Generic SSO with proper state priority handling.
3150 Optionally generates PKCE parameters if GENERIC_CLIENT_USE_PKCE is enabled.
3152 Priority order:
3153 1. CLI state (if provided)
3154 2. GENERIC_CLIENT_STATE environment variable
3155 3. Generated UUID (required by Okta and most OAuth providers)
3158 Args:
3159 state: Optional state parameter (e.g., CLI state)
3160 generic_authorization_endpoint: Authorization endpoint URL
3162 Returns:
3163 Tuple[dict, Optional[str]]:
3164 - Redirect parameters for SSO login (may include PKCE params)
3165 - code_verifier (if PKCE is enabled, None otherwise)
3166 """
3167 redirect_params: Final = {}
3168 code_verifier: str | None = None
3170 if state:
3171 # CLI state takes priority
3172 # the litellm proxy cli sends the "state" parameter to the proxy server for auth. We should maintain the state parameter for the cli if it is provided
3173 redirect_params["state"] = state
3174 else:
3175 generic_client_state: Final = os.getenv("GENERIC_CLIENT_STATE", None)
3176 if generic_client_state:
3177 redirect_params["state"] = generic_client_state
3178 else:
3179 redirect_params["state"] = uuid.uuid4().hex
3181 # Handle PKCE (Proof Key for Code Exchange) if enabled
3182 # Set GENERIC_CLIENT_USE_PKCE=true to enable PKCE for enhanced OAuth security
3183 use_pkce: Final = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
3185 if use_pkce:
3186 (
3187 code_verifier,
3188 code_challenge,
3189 ) = SSOAuthenticationHandler.generate_pkce_params()
3190 redirect_params["code_challenge"] = code_challenge
3191 redirect_params["code_challenge_method"] = "S256"
3192 verbose_proxy_logger.debug("PKCE enabled for authorization request")
3194 return redirect_params, code_verifier
3196 @staticmethod
3197 def should_use_sso_handler(
3198 google_client_id: str | None = None,
3199 microsoft_client_id: str | None = None,
3200 generic_client_id: str | None = None,
3201 ) -> bool:
3202 if google_client_id is not None or microsoft_client_id is not None or generic_client_id is not None:
3203 return True
3204 return False
3206 @staticmethod
3207 def get_redirect_url_for_sso(
3208 request: Request,
3209 sso_callback_route: str,
3210 existing_key: str | None = None,
3211 ) -> str:
3212 """
3213 Get the redirect URL for SSO
3215 Note: existing_key is not added to the URL to avoid changing the callback URL.
3216 It should be passed via the state parameter instead.
3217 """
3218 from litellm.proxy.utils import get_custom_url
3220 redirect_url = get_custom_url(request_base_url=str(request.base_url))
3221 if redirect_url.endswith("/"):
3222 redirect_url += sso_callback_route
3223 else:
3224 redirect_url += "/" + sso_callback_route
3226 return redirect_url
3228 @staticmethod
3229 async def upsert_sso_user(
3230 result: CustomOpenID | OpenID | dict | None,
3231 user_info: NewUserResponse | LiteLLM_UserTable | None,
3232 user_email: str | None,
3233 user_defined_values: SSOUserDefinedValues | None,
3234 prisma_client: PrismaClient,
3235 ):
3236 """
3237 Connects the SSO Users to the User Table in LiteLLM DB
3239 - If user on LiteLLM DB, update the user_email and user_role (if SSO provides valid role) with the SSO values
3240 - If user not on LiteLLM DB, insert the user into LiteLLM DB
3241 """
3242 try:
3243 if user_info is not None:
3244 user_id: Final = user_info.user_id
3245 update_data: Final = _build_sso_user_update_data(
3246 result=result,
3247 user_email=user_email,
3248 user_id=user_id,
3249 )
3251 await _user_meta_db(UserRepository(prisma_client)).update_many(
3252 where={"user_id": user_id}, data=update_data
3253 )
3254 else:
3255 verbose_proxy_logger.info("user not in DB, inserting user into LiteLLM DB")
3256 # user not in DB, insert User into LiteLLM DB
3257 user_info = await insert_sso_user(
3258 result_openid=result,
3259 user_defined_values=user_defined_values,
3260 )
3261 return user_info
3262 except Exception as e:
3263 verbose_proxy_logger.exception("Error upserting SSO user into LiteLLM DB: %s", e)
3264 return user_info
3266 @staticmethod
3267 async def add_user_to_teams_from_sso_response(
3268 result: CustomOpenID | OpenID | dict | None,
3269 user_info: NewUserResponse | LiteLLM_UserTable | None,
3270 ):
3271 """
3272 Adds the user as a team member to the teams specified in the SSO responses `team_ids` field
3275 The `team_ids` field is populated by litellm after processing the SSO response
3276 """
3277 if user_info is None:
3278 verbose_proxy_logger.debug("User not found in LiteLLM DB, skipping team member addition")
3279 return
3280 sso_teams: Final = getattr(result, "team_ids", [])
3281 await add_missing_team_member(user_info=user_info, sso_teams=sso_teams)
3283 @staticmethod
3284 def verify_user_in_restricted_sso_group(
3285 general_settings: dict,
3286 result: CustomOpenID | OpenID | dict | None,
3287 received_response: dict | None,
3288 ) -> Literal[True]:
3289 """
3290 when ui_access_mode.type == "restricted_sso_group":
3292 - result.team_ids should contain the restricted_sso_group
3293 - if not, raise a ProxyException
3294 - if so, return True
3295 - if result.team_ids is None, return False
3296 - if result.team_ids is an empty list, return False
3297 - if result.team_ids is a list, return True if the restricted_sso_group is in the list, otherwise return False
3298 """
3300 ui_access_mode: Final = cast(dict | str | None, general_settings.get("ui_access_mode"))
3302 if ui_access_mode is None:
3303 return True
3304 if isinstance(ui_access_mode, str):
3305 return True
3306 team_ids: Final = getattr(result, "team_ids", [])
3308 if ui_access_mode.get("type") == "restricted_sso_group":
3309 restricted_sso_group: Final = ui_access_mode.get("restricted_sso_group")
3310 if restricted_sso_group not in team_ids:
3311 raise ProxyException(
3312 message=f"User is not in the restricted SSO group: {restricted_sso_group}. User groups: {team_ids}. Received SSO response: {received_response}",
3313 type=ProxyErrorTypes.auth_error,
3314 param="restricted_sso_group",
3315 code=status.HTTP_403_FORBIDDEN,
3316 )
3317 return True
3319 @staticmethod
3320 async def create_litellm_team_from_sso_group(
3321 litellm_team_id: str,
3322 litellm_team_name: str | None = None,
3323 ):
3324 """
3325 Creates a Litellm Team from a SSO Group ID
3327 Your SSO provider might have groups that should be created on LiteLLM
3329 Use this helper to create a Litellm Team from a SSO Group ID
3331 Args:
3332 litellm_team_id (str): The ID of the Litellm Team
3333 litellm_team_name (Optional[str]): The name of the Litellm Team
3334 """
3335 from litellm.proxy.proxy_server import prisma_client
3337 if prisma_client is None:
3338 raise ProxyException(
3339 message="Prisma client not found. Set it in the proxy_server.py file",
3340 type=ProxyErrorTypes.auth_error,
3341 param="prisma_client",
3342 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
3343 )
3344 try:
3345 team_obj: Final = await _team_detail_db(TeamRepository(prisma_client)).find_first(
3346 where={"team_id": litellm_team_id}
3347 )
3348 verbose_proxy_logger.debug("Team object: %s", team_obj)
3350 # only create a new team if it doesn't exist
3351 if team_obj:
3352 verbose_proxy_logger.debug("Team already exists: %s - %s", litellm_team_id, litellm_team_name)
3353 return
3355 team_request: NewTeamRequest = NewTeamRequest(
3356 team_id=litellm_team_id,
3357 team_alias=litellm_team_name,
3358 )
3359 if litellm.default_team_params:
3360 team_request = SSOAuthenticationHandler._cast_and_deepcopy_litellm_default_team_params(
3361 default_team_params=litellm.default_team_params,
3362 litellm_team_id=litellm_team_id,
3363 litellm_team_name=litellm_team_name,
3364 team_request=team_request,
3365 )
3367 await new_team(
3368 data=team_request,
3369 # params used for Audit Logging
3370 http_request=Request(scope={"type": "http", "method": "POST"}),
3371 user_api_key_dict=UserAPIKeyAuth(
3372 token="",
3373 key_alias=f"litellm.{MicrosoftSSOHandler.__name__}",
3374 ),
3375 )
3376 except Exception as e:
3377 verbose_proxy_logger.exception("Error creating Litellm Team: %s", e)
3379 @staticmethod
3380 def _cast_and_deepcopy_litellm_default_team_params(
3381 default_team_params: DefaultTeamSSOParams | dict,
3382 team_request: NewTeamRequest,
3383 litellm_team_id: str,
3384 litellm_team_name: str | None = None,
3385 ) -> NewTeamRequest:
3386 """
3387 Casts and deepcopies the litellm.default_team_params to a NewTeamRequest object
3389 - Ensures we create a new DefaultTeamSSOParams object
3390 - Handle the case where litellm.default_team_params is a dict or a DefaultTeamSSOParams object
3391 - Adds the litellm_team_id and litellm_team_name to the DefaultTeamSSOParams object
3392 """
3393 if isinstance(default_team_params, dict):
3394 _team_request: Final = deepcopy(default_team_params)
3395 _team_request["team_id"] = litellm_team_id
3396 _team_request["team_alias"] = litellm_team_name
3397 team_request = NewTeamRequest(**_team_request)
3398 elif isinstance(litellm.default_team_params, DefaultTeamSSOParams):
3399 _default_team_params: Final = deepcopy(litellm.default_team_params)
3400 _new_team_request: Final = team_request.model_dump()
3401 _new_team_request.update(_default_team_params)
3402 team_request = NewTeamRequest.model_validate(_new_team_request)
3403 return team_request
3405 @staticmethod
3406 def _get_cli_state(
3407 source: str | None,
3408 key: str | None,
3409 existing_key: str | None = None,
3410 user_code: str | None = None,
3411 ) -> str | None:
3412 """
3413 Checks the request 'source' if a cli state token was passed in
3415 This is used to authenticate through the CLI login flow.
3417 The state parameter format is: {PREFIX}:{login_id}[:{user_code}]
3418 - The state parameter is used to pass data through the OAuth flow without changing the callback URL
3419 - user_code is appended only for the opt-in verification_uri_complete flow so the verify page can pre-fill it
3420 """
3421 from litellm.constants import (
3422 LITELLM_CLI_SESSION_TOKEN_PREFIX,
3423 )
3425 if source == LITELLM_CLI_SOURCE_IDENTIFIER and key:
3426 if _is_valid_cli_sso_user_code(user_code):
3427 return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}:{user_code}"
3428 return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}"
3429 else:
3430 return None
3432 @staticmethod
3433 def _get_user_email_and_id_from_result(
3434 result: OpenID | dict | None,
3435 generic_client_id: str | None = None,
3436 ) -> ParsedOpenIDResult:
3437 """
3438 Gets the user email and id from the OpenID result after validating the email domain
3439 """
3440 user_email: str | None = normalize_email(getattr(result, "email", None))
3441 user_id: str | None = getattr(result, "id", None) if result is not None else None
3442 user_role: str | None = None
3444 if user_email is not None and os.getenv("ALLOWED_EMAIL_DOMAINS") is not None:
3445 email_domain: Final = user_email.split("@")[1]
3446 allowed_domains: Final = os.getenv("ALLOWED_EMAIL_DOMAINS").split(",")
3447 if email_domain not in allowed_domains:
3448 raise HTTPException(
3449 status_code=401,
3450 detail={
3451 "message": f"The email domain={email_domain}, is not an allowed email domain={allowed_domains}. Contact your admin to change this."
3452 },
3453 )
3455 # Extract user_role from result (works for all SSO providers)
3456 if result is not None:
3457 _user_role: Final = getattr(result, "user_role", None)
3458 if _user_role is not None:
3459 # Convert enum to string if needed
3460 user_role = _user_role.value if isinstance(_user_role, LitellmUserRoles) else _user_role
3461 verbose_proxy_logger.debug("Extracted user_role from SSO result: %s", user_role)
3463 # generic client id - override with custom attribute name if specified
3464 if generic_client_id is not None and result is not None:
3465 generic_user_role_attribute_name: Final = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role")
3466 user_id = getattr(result, "id", None)
3467 user_email = normalize_email(getattr(result, "email", None))
3468 if user_role is None:
3469 _role_from_attr: Final = getattr(result, generic_user_role_attribute_name, None)
3470 if _role_from_attr is not None:
3471 # Convert enum to string if needed
3472 user_role = (
3473 _role_from_attr.value if isinstance(_role_from_attr, LitellmUserRoles) else _role_from_attr
3474 )
3476 if user_id is None and result is not None:
3477 _first_name: Final = getattr(result, "first_name", "") or ""
3478 _last_name: Final = getattr(result, "last_name", "") or ""
3479 user_id = _first_name + _last_name
3481 if user_email is not None and (user_id is None or len(user_id) == 0):
3482 user_id = user_email
3484 return ParsedOpenIDResult(
3485 user_email=user_email,
3486 user_id=user_id,
3487 user_role=user_role,
3488 )
3490 @staticmethod
3491 async def get_redirect_response_from_openid(
3492 result: OpenID | dict | CustomOpenID,
3493 request: Request,
3494 received_response: dict | None = None,
3495 generic_client_id: str | None = None,
3496 ui_access_mode: dict | None = None,
3497 access_token_payload: dict | None = None,
3498 jwt_handler: JWTHandler | None = None,
3499 return_to: str | None = None,
3500 sso_assertion: SSOIdentityAssertion | None = None,
3501 ) -> RedirectResponse:
3503 from litellm.proxy.proxy_server import (
3504 general_settings,
3505 generate_key_helper_fn,
3506 master_key,
3507 premium_user,
3508 proxy_logging_obj,
3509 redis_usage_cache,
3510 user_api_key_cache,
3511 user_custom_sso,
3512 )
3513 from litellm.proxy.utils import get_prisma_client_or_throw
3514 from litellm.types.proxy.ui_sso import ReturnedUITokenObject
3516 prisma_client: Final = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy")
3518 # User is Authe'd in - generate key for the UI to access Proxy
3519 parsed_openid_result: Final = SSOAuthenticationHandler._get_user_email_and_id_from_result(
3520 result=result, generic_client_id=generic_client_id
3521 )
3522 user_email: Final = parsed_openid_result.get("user_email")
3523 user_id = parsed_openid_result.get("user_id")
3524 user_role = parsed_openid_result.get("user_role")
3525 verbose_proxy_logger.info("SSO callback result: %s", result)
3527 user_info = None
3528 user_id_models: Final[list] = []
3529 max_internal_user_budget: Final = litellm.max_internal_user_budget
3530 internal_user_budget_duration: Final = litellm.internal_user_budget_duration
3532 # User might not be already created on first generation of key
3533 # But if it is, we want their models preferences
3534 user_defined_values: SSOUserDefinedValues | None = None
3536 custom_sso_handler: Final[_CustomSsoCall | None] = user_custom_sso
3537 if custom_sso_handler is not None:
3538 if inspect.iscoroutinefunction(custom_sso_handler):
3539 user_defined_values = await custom_sso_handler(result)
3540 else:
3541 raise ValueError("user_custom_sso must be a coroutine function")
3542 elif user_id is not None:
3543 user_defined_values = SSOUserDefinedValues(
3544 models=user_id_models,
3545 user_id=user_id,
3546 user_email=user_email,
3547 max_budget=max_internal_user_budget,
3548 user_role=user_role,
3549 budget_duration=internal_user_budget_duration,
3550 )
3552 # (IF SET) Verify user is in restricted SSO group
3553 SSOAuthenticationHandler.verify_user_in_restricted_sso_group(
3554 general_settings=general_settings,
3555 result=result,
3556 received_response=received_response,
3557 )
3559 user_info = await get_user_info_from_db(
3560 result=result,
3561 prisma_client=prisma_client,
3562 user_api_key_cache=user_api_key_cache,
3563 proxy_logging_obj=proxy_logging_obj,
3564 user_email=user_email,
3565 user_defined_values=user_defined_values,
3566 alternate_user_id=user_id,
3567 )
3569 # Sync user role from JWT claims via jwt_litellm_role_map (if configured).
3570 # This ensures SSO users get the same role mapping as API/JWT users.
3571 # Use the decoded access_token_payload (not received_response) because
3572 # custom role claims (e.g. custom_roles) are encoded inside the JWT
3573 # access token, which is stripped from received_response.
3574 await _sync_user_role_from_jwt_role_map(
3575 jwt_handler=jwt_handler,
3576 received_response=access_token_payload or received_response,
3577 user_info=user_info,
3578 prisma_client=prisma_client,
3579 user_api_key_cache=user_api_key_cache,
3580 user_defined_values=user_defined_values,
3581 )
3583 user_defined_values = apply_user_info_values_to_sso_user_defined_values(
3584 user_info=user_info, user_defined_values=user_defined_values
3585 )
3587 if user_defined_values is None:
3588 raise Exception(
3589 "Unable to map user identity to known values. 'user_defined_values' is None. File an issue - https://github.com/BerriAI/litellm/issues"
3590 )
3592 verbose_proxy_logger.info("user_defined_values for creating ui key: %s", user_defined_values)
3594 response: Final = await generate_key_helper_fn(
3595 llm_router=None,
3596 request_type="key",
3597 duration=LITELLM_UI_SESSION_DURATION,
3598 key_max_budget=litellm.max_ui_session_budget,
3599 aliases={},
3600 config={},
3601 spend=0,
3602 team_id="litellm-dashboard",
3603 models=user_defined_values["models"],
3604 user_id=user_defined_values["user_id"],
3605 user_email=user_defined_values["user_email"],
3606 user_role=user_defined_values["user_role"],
3607 max_budget=user_defined_values["max_budget"],
3608 budget_duration=user_defined_values["budget_duration"],
3609 table_name="key",
3610 )
3612 key = response["token"]
3613 user_id = response["user_id"]
3615 user_role = user_defined_values["user_role"] or LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value
3616 if user_id and isinstance(user_id, str):
3617 user_role = await check_and_update_if_proxy_admin_id(
3618 user_role=user_role, user_id=user_id, prisma_client=prisma_client
3619 )
3621 verbose_proxy_logger.debug("user_role: %s; ui_access_mode: %s", user_role, ui_access_mode)
3622 ## CHECK IF ROLE ALLOWED TO USE PROXY ##
3623 is_admin_only_access: Final = check_is_admin_only_access(ui_access_mode or {})
3624 if is_admin_only_access:
3625 has_access: Final = has_admin_ui_access(user_role or "")
3626 if not has_access:
3627 raise HTTPException(
3628 status_code=401,
3629 detail={
3630 "error": f"User not allowed to access proxy. User role={user_role}, proxy mode={ui_access_mode}"
3631 },
3632 )
3634 if isinstance(user_id, str) and user_id:
3635 await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
3636 await warn_if_id_jag_assertion_uncaptured(sso_assertion)
3638 disabled_non_admin_personal_key_creation: Final = get_disabled_non_admin_personal_key_creation()
3639 litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/")
3641 if get_secret_bool("EXPERIMENTAL_UI_LOGIN"):
3642 _user_info: LiteLLM_UserTable | None = None
3643 if user_defined_values is not None and user_defined_values["user_id"] is not None:
3644 _user_info = LiteLLM_UserTable(
3645 user_id=user_defined_values["user_id"],
3646 user_role=user_defined_values["user_role"] or user_role,
3647 models=[],
3648 max_budget=litellm.max_ui_session_budget,
3649 )
3650 if _user_info is None:
3651 raise HTTPException(
3652 status_code=401,
3653 detail={"error": "User Information is required for experimental UI login"},
3654 )
3656 key = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(_user_info)
3658 returned_ui_token_object: Final = ReturnedUITokenObject(
3659 user_id=cast(str, user_id),
3660 key=key,
3661 user_email=user_email,
3662 user_role=user_role or LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
3663 login_method="sso",
3664 premium_user=premium_user,
3665 auth_header_name=general_settings.get("litellm_key_header_name", "Authorization"),
3666 disabled_non_admin_personal_key_creation=disabled_non_admin_personal_key_creation,
3667 server_root_path=get_server_root_path(),
3668 password_reset_required=False,
3669 )
3671 from litellm.proxy.auth.login_utils import encode_ui_session_jwt
3673 jwt_token: Final = encode_ui_session_jwt(returned_ui_token_object, master_key or "")
3675 # Post-SSO return_to handling (the same-origin DCR round-trip and the control-plane
3676 # cross-origin code exchange) lives in one shared helper so this method stays inside the
3677 # complexity budget. None falls through to the dashboard redirect below.
3678 return_to_redirect: Final = await _sso_return_to_redirect(
3679 return_to=return_to,
3680 jwt_token=jwt_token,
3681 redis_usage_cache=redis_usage_cache,
3682 user_api_key_cache=user_api_key_cache,
3683 request=request,
3684 )
3685 if return_to_redirect is not None:
3686 return return_to_redirect
3688 if user_id is not None and isinstance(user_id, str):
3689 litellm_dashboard_ui += "?login=success"
3690 verbose_proxy_logger.info("Redirecting to %s", litellm_dashboard_ui)
3691 redirect_response: Final = RedirectResponse(url=litellm_dashboard_ui, status_code=303)
3692 set_session_token_cookie(redirect_response, request, jwt_token)
3693 return redirect_response
3695 @staticmethod
3696 async def prepare_token_exchange_parameters(
3697 request: Request,
3698 generic_include_client_id: bool,
3699 ) -> dict:
3700 """
3701 Prepare token exchange parameters for Generic SSO.
3703 Args:
3704 request: Request object
3705 generic_include_client_id: Generic OAuth Client ID
3707 Returns:
3708 dict: Token exchange parameters
3709 """
3710 # Prepare token exchange parameters (may add code_verifier: str later)
3711 token_params: Final[dict[str, object]] = {"include_client_id": generic_include_client_id}
3713 # Retrieve PKCE code_verifier if PKCE was used in authorization.
3714 # Gate on GENERIC_CLIENT_USE_PKCE to avoid an unnecessary Redis round-trip
3715 # on every non-PKCE SSO callback.
3716 query_params: Final = dict(request.query_params)
3717 state: Final = query_params.get("state")
3719 use_pkce: Final = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
3721 if use_pkce and not state:
3722 verbose_proxy_logger.warning(
3723 "PKCE is enabled (GENERIC_CLIENT_USE_PKCE=true) but no 'state' parameter "
3724 "was found in the callback. The PKCE verifier cannot be retrieved without "
3725 "a state value — the token exchange will proceed without code_verifier, "
3726 "which the provider may reject. Ensure your OAuth provider returns 'state' "
3727 "in the callback redirect."
3728 )
3730 if state and use_pkce:
3731 from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
3733 cache_key: Final = f"pkce_verifier:{state}"
3734 if redis_usage_cache is not None:
3735 cached_data = await redis_usage_cache.async_get_cache(key=cache_key)
3736 else:
3737 cached_data = await user_api_key_cache.async_get_cache(key=cache_key)
3739 code_verifier = None
3740 # Track why code_verifier is absent for accurate strict-mode diagnostics.
3741 _empty_value_in_dict = False # dict format correct but value is empty/null
3743 if cached_data:
3744 # Extract code_verifier from dict (stored as dict for JSON serialization)
3745 if isinstance(cached_data, dict) and "code_verifier" in cached_data:
3746 code_verifier = cached_data["code_verifier"]
3747 if not code_verifier:
3748 # Dict format is correct but value is empty or null. This is
3749 # a distinct case from an unrecognized format — the entry exists
3750 # but was stored with an empty/null verifier (data integrity issue).
3751 _empty_value_in_dict = True
3752 verbose_proxy_logger.warning(
3753 "PKCE verifier dict for state '%s' has an empty/null code_verifier "
3754 "value — may indicate a storage bug. Treating as a cache miss.",
3755 state,
3756 )
3757 else:
3758 verbose_proxy_logger.debug("PKCE code_verifier retrieved from cache")
3759 elif isinstance(cached_data, str):
3760 # Handle legacy format (plain string) for backward compatibility
3761 code_verifier = cached_data
3762 verbose_proxy_logger.warning(
3763 "Retrieved code_verifier in legacy plain-string format. Future storage will use dict format."
3764 )
3765 else:
3766 # Defer the detailed ERROR log to the strict-mode branch below
3767 # (which includes state and a diagnostic message). Log at DEBUG
3768 # here to avoid duplicate ERROR entries in the same request.
3769 verbose_proxy_logger.debug(
3770 "Unexpected PKCE verifier cache format (type=%s); skipping.",
3771 type(cached_data).__name__,
3772 )
3774 if code_verifier:
3775 # Add code_verifier to token exchange parameters.
3776 token_params["code_verifier"] = code_verifier
3777 # Return the cache key so the caller can delete it *after* a
3778 # successful token exchange (avoids losing the verifier on retry
3779 # if the exchange fails partway through).
3780 token_params["_pkce_cache_key"] = cache_key
3781 else:
3782 await SSOAuthenticationHandler._handle_missing_pkce_verifier(
3783 state=state,
3784 cache_key=cache_key,
3785 cached_data=cached_data,
3786 empty_value_in_dict=_empty_value_in_dict,
3787 redis_usage_cache=redis_usage_cache,
3788 user_api_key_cache=user_api_key_cache,
3789 )
3790 return token_params
3792 @staticmethod
3793 async def _handle_missing_pkce_verifier(
3794 state: str | None,
3795 cache_key: str,
3796 cached_data: object,
3797 empty_value_in_dict: bool,
3798 redis_usage_cache: object,
3799 user_api_key_cache: object,
3800 ) -> None:
3801 """Handle the case where PKCE verifier could not be extracted from cache.
3803 In strict mode (PKCE_STRICT_CACHE_MISS=true) raises ProxyException.
3804 Otherwise logs a warning and returns (token exchange proceeds without verifier).
3805 """
3806 active_cache: Final = redis_usage_cache if redis_usage_cache is not None else user_api_key_cache
3807 strict_cache_miss: Final = os.getenv("PKCE_STRICT_CACHE_MISS", "false").lower() == "true"
3808 if strict_cache_miss:
3809 if empty_value_in_dict:
3810 await SSOAuthenticationHandler._delete_pkce_verifier(cache_key)
3811 raise ProxyException(
3812 message=(
3813 f"PKCE verifier for state '{state}' was found in cache but "
3814 f"has an empty or null code_verifier value — possible storage bug."
3815 ),
3816 type=ProxyErrorTypes.auth_error,
3817 param="PKCE_CACHE_MISS",
3818 code=status.HTTP_401_UNAUTHORIZED,
3819 )
3820 elif cached_data is not None:
3821 await SSOAuthenticationHandler._delete_pkce_verifier(cache_key)
3822 verbose_proxy_logger.error(
3823 "PKCE verifier for state '%s' has an unrecognized format (type=%s); "
3824 "treating as a cache miss. Investigate the cached value — it may be "
3825 "a corrupt or stale entry.",
3826 state,
3827 type(cached_data).__name__,
3828 )
3829 raise ProxyException(
3830 message=(
3831 f"PKCE verifier for state '{state}' has an unrecognized format "
3832 f"(type={type(cached_data).__name__}). The cached entry may be corrupt."
3833 ),
3834 type=ProxyErrorTypes.auth_error,
3835 param="PKCE_CACHE_MISS",
3836 code=status.HTTP_401_UNAUTHORIZED,
3837 )
3838 else:
3839 if redis_usage_cache is not None:
3840 cause = (
3841 "The authorization and callback were likely handled by different "
3842 "instances — the verifier was stored on one pod but not found on another."
3843 )
3844 else:
3845 cause = (
3846 "The verifier may have expired (TTL), been lost on a pod restart, "
3847 "or the PKCE authorization step was never completed. "
3848 "Configure Redis so all proxy instances share the PKCE verifier."
3849 )
3850 verbose_proxy_logger.error(
3851 "PKCE is enabled but no verifier found in cache for state '%s'. %s Cache type: %s.",
3852 state,
3853 cause,
3854 type(active_cache).__name__,
3855 )
3856 raise ProxyException(
3857 message=f"PKCE verifier not found in cache for state '{state}'. {cause}",
3858 type=ProxyErrorTypes.auth_error,
3859 param="PKCE_CACHE_MISS",
3860 code=status.HTTP_401_UNAUTHORIZED,
3861 )
3862 else:
3863 if cached_data is not None:
3864 await SSOAuthenticationHandler._delete_pkce_verifier(cache_key)
3865 verbose_proxy_logger.warning(
3866 "PKCE is enabled but verifier not found in cache for state '%s' "
3867 "(cache type: %s, raw data present: %s). "
3868 "Continuing without code_verifier — set PKCE_STRICT_CACHE_MISS=true to fail fast instead.",
3869 state,
3870 type(active_cache).__name__,
3871 cached_data is not None,
3872 )
3874 @staticmethod
3875 async def _delete_pkce_verifier(cache_key: str) -> None:
3876 """Delete a single-use PKCE verifier from cache after a successful exchange.
3878 Failure is non-fatal: a leftover verifier is a minor security concern
3879 (unused key in cache) but not worth aborting an otherwise-successful login.
3880 """
3881 from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
3883 try:
3884 if redis_usage_cache is not None:
3885 await redis_usage_cache.async_delete_cache(key=cache_key)
3886 else:
3887 await user_api_key_cache.async_delete_cache(key=cache_key)
3888 except Exception as exc:
3889 verbose_proxy_logger.warning(
3890 "PKCE: failed to delete verifier cache key '%s' (best-effort cleanup): %s",
3891 cache_key,
3892 exc,
3893 )
3895 @staticmethod
3896 def generate_pkce_params() -> tuple[str, str]:
3897 """
3898 Generate PKCE (Proof Key for Code Exchange) parameters for OAuth 2.0.
3900 Returns:
3901 Tuple[str, str]: (code_verifier, code_challenge)
3902 - code_verifier: Random 43-128 character string (we use 43 for efficiency)
3903 - code_challenge: Base64-URL-encoded SHA256 hash of the code_verifier
3905 Reference: https://datatracker.ietf.org/doc/html/rfc7636
3906 """
3907 # Generate a cryptographically random code_verifier (43 characters)
3908 # Using 32 random bytes which becomes 43 characters when base64-url-encoded
3909 code_verifier: Final = base64.urlsafe_b64encode(secrets.token_bytes(32)).decode("utf-8").rstrip("=")
3911 # Generate code_challenge using S256 method (SHA256)
3912 code_challenge_bytes: Final = hashlib.sha256(code_verifier.encode("utf-8")).digest()
3913 code_challenge: Final = base64.urlsafe_b64encode(code_challenge_bytes).decode("utf-8").rstrip("=")
3915 return code_verifier, code_challenge
3917 @staticmethod
3918 def _validate_token_response(response: "httpx.Response") -> dict:
3919 """
3920 Parse and validate the token endpoint response.
3922 Ensures the response is valid JSON, a dict, and contains a non-null
3923 access_token string. Raises ProxyException on any validation failure.
3924 """
3925 try:
3926 token_response_raw: Final[object] = _as_object(response.json())
3927 except Exception as json_err:
3928 verbose_proxy_logger.error(
3929 "Failed to parse token response as JSON: %s. Body: %s",
3930 json_err,
3931 response.text[:500],
3932 )
3933 raise ProxyException(
3934 message=f"Token endpoint returned invalid JSON: {json_err}",
3935 type=ProxyErrorTypes.auth_error,
3936 param="token_exchange",
3937 code=status.HTTP_401_UNAUTHORIZED,
3938 )
3940 if not isinstance(token_response_raw, dict):
3941 verbose_proxy_logger.error(
3942 "Token endpoint returned non-dict JSON (type=%s). Body: %s",
3943 type(token_response_raw).__name__,
3944 response.text[:500],
3945 )
3946 raise ProxyException(
3947 message=(
3948 f"Token endpoint returned unexpected response format "
3949 f"(expected JSON object, got {type(token_response_raw).__name__})"
3950 ),
3951 type=ProxyErrorTypes.auth_error,
3952 param="token_exchange",
3953 code=status.HTTP_401_UNAUTHORIZED,
3954 )
3955 token_response: Final[dict] = token_response_raw
3957 access_token_val: Final = token_response.get("access_token")
3958 if not isinstance(access_token_val, str) or not access_token_val:
3959 error: Final = token_response.get("error")
3960 error_desc: Final = token_response.get("error_description", "")
3961 if error:
3962 detail = f"{error} - {error_desc}" if error_desc else error
3963 else:
3964 detail = (
3965 "token endpoint returned HTTP 200 but no access_token "
3966 f"(response keys: {sorted(token_response.keys())})"
3967 )
3968 verbose_proxy_logger.error("Token response missing or null access_token. detail=%s", detail)
3969 raise ProxyException(
3970 message=f"Token exchange failed: {detail}",
3971 type=ProxyErrorTypes.auth_error,
3972 param="token_exchange",
3973 code=status.HTTP_401_UNAUTHORIZED,
3974 )
3976 return token_response
3978 @staticmethod
3979 async def _pkce_token_exchange(
3980 authorization_code: str,
3981 code_verifier: str,
3982 client_id: str,
3983 client_secret: str | None,
3984 token_endpoint: str,
3985 userinfo_endpoint: str | None,
3986 include_client_id: bool,
3987 redirect_url: str | None,
3988 additional_headers: dict[str, str],
3989 ) -> dict:
3990 """
3991 Performs a direct OAuth token exchange including the PKCE code_verifier.
3993 fastapi-sso does not forward code_verifier, so when PKCE is enabled we
3994 bypass it and call the token endpoint ourselves, then fetch user info.
3996 Returns a combined dict of the token response and user info, suitable
3997 for passing to a response_convertor.
3998 """
3999 verbose_proxy_logger.debug(
4000 "PKCE: performing direct token exchange (code_verifier length=%d)",
4001 len(code_verifier),
4002 )
4004 token_data: Final[dict[str, str]] = {
4005 "grant_type": "authorization_code",
4006 "code": authorization_code,
4007 "code_verifier": code_verifier,
4008 }
4009 # Only include redirect_uri when set — omitting it avoids sending the
4010 # literal string "None" to the provider if the env var is missing.
4011 if redirect_url:
4012 token_data["redirect_uri"] = redirect_url
4014 request_headers: Final = {
4015 **additional_headers,
4016 "Content-Type": "application/x-www-form-urlencoded", # must not be overridden
4017 "Accept": "application/json",
4018 }
4020 if not include_client_id:
4021 # Use Basic Auth only when a secret is available; public PKCE clients omit it.
4022 if client_secret:
4023 credentials: Final = base64.b64encode(f"{client_id}:{client_secret}".encode()).decode()
4024 request_headers["Authorization"] = f"Basic {credentials}"
4025 else:
4026 token_data["client_id"] = client_id
4027 else:
4028 token_data["client_id"] = client_id
4029 if client_secret:
4030 token_data["client_secret"] = client_secret
4032 http_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER)
4033 try:
4034 response: Final = await http_client.post(
4035 url=token_endpoint,
4036 data=token_data,
4037 headers=request_headers,
4038 timeout=30.0,
4039 )
4040 except Exception as exc:
4041 # Catch network-level errors (SSL, DNS, TCP, timeout, etc.) and
4042 # wrap them as a clean ProxyException rather than leaking raw
4043 # httpx or OS exceptions to callers.
4044 verbose_proxy_logger.error("PKCE token endpoint unreachable: %s", exc)
4045 raise ProxyException(
4046 message=f"Token endpoint request failed: {exc}",
4047 type=ProxyErrorTypes.auth_error,
4048 param="token_exchange",
4049 code=status.HTTP_401_UNAUTHORIZED,
4050 ) from exc
4051 if response.status_code != 200:
4052 verbose_proxy_logger.error(
4053 "PKCE token exchange failed. status=%s body=%s",
4054 response.status_code,
4055 response.text[:500],
4056 )
4057 raise ProxyException(
4058 message=f"Token exchange failed: {response.status_code} - {response.text[:500]}",
4059 type=ProxyErrorTypes.auth_error,
4060 param="token_exchange",
4061 code=status.HTTP_401_UNAUTHORIZED,
4062 )
4064 token_response: Final = SSOAuthenticationHandler._validate_token_response(response)
4066 verbose_proxy_logger.debug(
4067 "PKCE token exchange successful. id_token_present=%s",
4068 bool(token_response.get("id_token")),
4069 )
4070 # Bearer credentials (access_token, id_token, refresh_token) are always sourced
4071 # from token_response — not from userinfo — in the merge step below.
4072 userinfo: Final = await SSOAuthenticationHandler._get_pkce_userinfo(
4073 access_token=token_response["access_token"],
4074 id_token=token_response.get("id_token"),
4075 userinfo_endpoint=userinfo_endpoint,
4076 additional_headers=additional_headers,
4077 )
4079 # Merge: userinfo takes precedence for identity claims (sub, email, name, …) per
4080 # the OpenID Connect spec (userinfo is the authoritative source for identity).
4081 # Bearer credentials (access_token, id_token, refresh_token) from the token endpoint
4082 # take precedence over same-named fields in userinfo — non-standard providers sometimes
4083 # include token fields in userinfo, which must not shadow the real bearer token.
4084 # If a bearer field is absent from the token response, any userinfo-provided value
4085 # is preserved as a fallback (useful for non-standard providers that omit id_token
4086 # from the token response but include it in userinfo).
4087 #
4088 # Three-way merge semantics for each bearer-credential field:
4089 # 1. token_response has a non-null value → use it (token endpoint is authoritative)
4090 # 2. token_response explicitly sent null → remove the key so callers get a clean
4091 # absence signal; the null from the token endpoint overrides userinfo too
4092 # 3. field absent from token_response → leave whatever userinfo provided as-is
4093 # (e.g. userinfo-provided id_token from a non-standard provider)
4094 merged: Final = {**token_response, **userinfo}
4095 for field in _OAUTH_TOKEN_FIELDS:
4096 if token_response.get(field) is not None:
4097 # Case 1: non-null in token_response — restore authoritative value.
4098 merged[field] = token_response[field]
4099 elif field in token_response:
4100 # Case 2: key exists but value is explicitly null — remove from merged.
4101 merged.pop(field, None)
4102 # Case 3: field absent from token_response — leave userinfo value as-is.
4103 return merged
4105 @staticmethod
4106 async def _get_pkce_userinfo(
4107 access_token: str,
4108 id_token: str | None,
4109 userinfo_endpoint: str | None,
4110 additional_headers: dict[str, str],
4111 ) -> dict:
4112 """
4113 Fetches user info from the userinfo endpoint.
4114 Falls back to decoding the id_token if the endpoint is unavailable.
4115 """
4116 # None = request not yet attempted, failed, or returned empty/null (treated as failure
4117 # so the id_token fallback can be attempted instead of returning a session with no claims).
4118 userinfo: dict | None = None
4120 if userinfo_endpoint:
4121 try:
4122 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER)
4123 resp: Final = await client.get(
4124 url=userinfo_endpoint,
4125 headers={
4126 **additional_headers,
4127 "Authorization": f"Bearer {access_token}", # must not be overridden
4128 },
4129 )
4130 if resp.status_code == 200:
4131 try:
4132 userinfo_raw: Final[dict[str, object] | None] = resp.json()
4133 if not userinfo_raw:
4134 # JSON null (None) or empty dict ({}) — no identity claims.
4135 # Treat as failure so id_token fallback can be attempted.
4136 verbose_proxy_logger.warning(
4137 "Userinfo endpoint returned an empty or null response "
4138 "(type=%s); treating as failure and attempting id_token fallback. "
4139 "Check your provider's userinfo endpoint configuration.",
4140 type(userinfo_raw).__name__,
4141 )
4142 userinfo = None
4143 else:
4144 userinfo = userinfo_raw
4145 except Exception as json_err:
4146 verbose_proxy_logger.warning(
4147 "Userinfo endpoint returned non-JSON response (status 200): %s",
4148 json_err,
4149 )
4150 else:
4151 verbose_proxy_logger.warning(
4152 "Userinfo endpoint returned %s (body: %s), falling back to id_token",
4153 resp.status_code,
4154 resp.text[:500],
4155 )
4156 except Exception as e:
4157 verbose_proxy_logger.warning("Userinfo endpoint error: %s, falling back to id_token", e)
4159 # Only fall back to id_token when the userinfo request failed (None).
4160 # Empty dict ({}) and JSON null are both treated as failure (set to None above) since
4161 # they contain no identity claims — id_token fallback is attempted in that case too.
4162 # Explicitly check for a non-empty string to avoid attempting JWT decode on
4163 # a blank or non-string id_token field from a misbehaving provider.
4164 if userinfo is None and isinstance(id_token, str) and id_token:
4165 try:
4166 userinfo = jwt.decode(id_token, options={"verify_signature": False})
4167 if not userinfo:
4168 # jwt.decode returned an empty dict (payload-free JWT or provider bug).
4169 # Treat this the same as a missing userinfo — the session would have no
4170 # identity claims, which is equivalent to a broken session.
4171 verbose_proxy_logger.warning("id_token decoded to an empty payload — treating as failure.")
4172 userinfo = None
4173 except Exception as decode_err:
4174 verbose_proxy_logger.error("Failed to decode id_token: %s", decode_err)
4175 raise ProxyException(
4176 message=f"Failed to decode id_token JWT: {decode_err}",
4177 type=ProxyErrorTypes.auth_error,
4178 param="userinfo",
4179 code=status.HTTP_401_UNAUTHORIZED,
4180 )
4182 if userinfo is None:
4183 id_token_attempted: Final = isinstance(id_token, str) and bool(id_token)
4184 if userinfo_endpoint:
4185 if id_token_attempted:
4186 detail = (
4187 "userinfo endpoint failed and id_token was present but "
4188 "decoded to an empty payload — no identity claims available"
4189 )
4190 else:
4191 detail = "userinfo endpoint failed and no id_token was present in the token response"
4192 else:
4193 if id_token_attempted:
4194 detail = (
4195 "no userinfo endpoint is configured (GENERIC_USERINFO_ENDPOINT) "
4196 "and id_token decoded to an empty payload — no identity claims available"
4197 )
4198 else:
4199 detail = (
4200 "no userinfo endpoint is configured (GENERIC_USERINFO_ENDPOINT) and no id_token was present"
4201 )
4202 raise ProxyException(
4203 message=f"SSO user info unavailable: {detail}.",
4204 type=ProxyErrorTypes.auth_error,
4205 param="userinfo",
4206 code=status.HTTP_401_UNAUTHORIZED,
4207 )
4209 return userinfo
4212class MicrosoftSSOHandler:
4213 """
4214 Handles Microsoft SSO callback response and returns a CustomOpenID object
4215 """
4217 DEFAULT_GRAPH_API_BASE_URL = "https://graph.microsoft.com/v1.0"
4219 """
4220 Constants
4221 """
4222 MAX_GRAPH_API_PAGES = 200
4224 # used for debugging to show the user groups litellm found from Graph API
4225 GRAPH_API_RESPONSE_KEY = "graph_api_user_groups"
4227 @staticmethod
4228 def get_graph_api_base_url() -> str:
4229 """
4230 Returns the Microsoft Graph API base URL, configurable via the
4231 `MICROSOFT_GRAPH_ENDPOINT` env var so non-default clouds such as Azure
4232 Government (GCC High) can point at `https://graph.microsoft.us/v1.0`
4233 """
4234 return get_secret_str("MICROSOFT_GRAPH_ENDPOINT") or MicrosoftSSOHandler.DEFAULT_GRAPH_API_BASE_URL
4236 @staticmethod
4237 def get_graph_api_user_groups_endpoint() -> str:
4238 return f"{MicrosoftSSOHandler.get_graph_api_base_url()}/me/memberOf"
4240 @staticmethod
4241 async def get_microsoft_callback_response(
4242 request: Request,
4243 microsoft_client_id: str,
4244 redirect_url: str,
4245 return_raw_sso_response: bool = False,
4246 ) -> CustomOpenID | OpenID | dict:
4247 """
4248 Get the Microsoft SSO callback response
4250 Args:
4251 return_raw_sso_response: If True, return the raw SSO response
4252 """
4253 microsoft_client_secret: Final = os.getenv("MICROSOFT_CLIENT_SECRET", None)
4254 microsoft_tenant: Final = os.getenv("MICROSOFT_TENANT", None)
4255 if microsoft_client_secret is None:
4256 raise ProxyException(
4257 message="MICROSOFT_CLIENT_SECRET not set. Set it in .env file",
4258 type=ProxyErrorTypes.auth_error,
4259 param="MICROSOFT_CLIENT_SECRET",
4260 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
4261 )
4262 if microsoft_tenant is None:
4263 raise ProxyException(
4264 message="MICROSOFT_TENANT not set. Set it in .env file",
4265 type=ProxyErrorTypes.auth_error,
4266 param="MICROSOFT_TENANT",
4267 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
4268 )
4269 microsoft_sso: Final = CustomMicrosoftSSO(
4270 client_id=microsoft_client_id,
4271 client_secret=microsoft_client_secret,
4272 tenant=microsoft_tenant,
4273 redirect_uri=redirect_url,
4274 allow_insecure_http=True,
4275 )
4276 original_msft_result: Final = (
4277 await microsoft_sso.verify_and_process(
4278 request=request,
4279 convert_response=False,
4280 )
4281 or {}
4282 )
4284 user_team_ids: Final = await MicrosoftSSOHandler.get_user_groups_from_graph_api(
4285 access_token=microsoft_sso.access_token
4286 )
4288 # Extract app roles from the id_token JWT
4289 app_roles: Final = MicrosoftSSOHandler.get_app_roles_from_id_token(id_token=microsoft_sso.id_token)
4290 verbose_proxy_logger.debug("Extracted app roles from id_token: %s", app_roles)
4292 # Combine groups and app roles
4293 user_role: Final = MicrosoftSSOHandler.get_user_role_from_app_roles(app_roles)
4295 verbose_proxy_logger.debug("Combined team_ids (groups + app roles): %s", user_team_ids)
4297 # if user is trying to get the raw sso response for debugging, return the raw sso response
4298 if return_raw_sso_response:
4299 original_msft_result[MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY] = user_team_ids
4300 original_msft_result["app_roles"] = app_roles
4301 return original_msft_result or {}
4303 result: Final = MicrosoftSSOHandler.openid_from_response(
4304 response=original_msft_result,
4305 team_ids=user_team_ids,
4306 user_role=user_role,
4307 )
4308 return result
4310 @staticmethod
4311 def openid_from_response(
4312 response: dict | None,
4313 team_ids: list[str],
4314 user_role: LitellmUserRoles | None,
4315 ) -> CustomOpenID:
4316 response = response or {}
4317 verbose_proxy_logger.debug("Microsoft SSO Callback Response: %s", response)
4318 openid_response: Final = CustomOpenID(
4319 email=normalize_email(response.get(MICROSOFT_USER_EMAIL_ATTRIBUTE) or response.get("mail")),
4320 display_name=response.get(MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE),
4321 provider="microsoft",
4322 id=response.get(MICROSOFT_USER_ID_ATTRIBUTE),
4323 first_name=response.get(MICROSOFT_USER_FIRST_NAME_ATTRIBUTE),
4324 last_name=response.get(MICROSOFT_USER_LAST_NAME_ATTRIBUTE),
4325 team_ids=team_ids,
4326 user_role=user_role,
4327 )
4328 verbose_proxy_logger.debug("Microsoft SSO OpenID Response: %s", openid_response)
4329 return openid_response
4331 @staticmethod
4332 def get_user_role_from_app_roles(
4333 app_roles: Sequence[str] | None,
4334 ) -> LitellmUserRoles | None:
4335 """
4336 Resolve the one role LiteLLM stores for a user from their Entra app roles.
4338 Entra does not guarantee `roles` claim ordering, so a user holding several app
4339 roles resolves to the highest privilege one rather than whichever the claim
4340 listed first. Roles the hierarchy does not rank (org_admin, team, customer)
4341 resolve by name to stay deterministic
4342 """
4343 return get_litellm_user_role(tuple(app_roles or ()))
4345 @staticmethod
4346 def get_app_roles_from_id_token(id_token: str | None) -> list[str]:
4347 """
4348 Extract app roles from the Microsoft Entra ID (Azure AD) id_token JWT.
4350 App roles are assigned in the Azure AD Enterprise Application and appear
4351 in the 'app_roles' claim of the id_token.
4353 Args:
4354 id_token (Optional[str]): The JWT id_token from Microsoft SSO
4356 Returns:
4357 List[str]: List of app role names assigned to the user
4358 """
4359 if not id_token:
4360 verbose_proxy_logger.debug("No id_token provided for app role extraction")
4361 return []
4363 try:
4364 import jwt
4366 # Decode the JWT without signature verification
4367 # (signature is already verified by fastapi_sso)
4368 decoded_token: Final = jwt.decode(id_token, options={"verify_signature": False})
4370 # Extract app_roles claim from the token
4371 ## check for both 'roles' and 'app_roles' claims
4372 roles: Final = decoded_token.get("app_roles", []) or decoded_token.get("roles", [])
4374 if roles and isinstance(roles, list):
4375 verbose_proxy_logger.debug("Found %s app role(s) in id_token: %s", len(roles), roles)
4376 return roles
4377 else:
4378 verbose_proxy_logger.debug("No app roles found in id_token or roles claim is not a list")
4379 return []
4381 except Exception as e:
4382 verbose_proxy_logger.error("Error extracting app roles from id_token: %s", e)
4383 return []
4385 @staticmethod
4386 async def get_user_groups_from_graph_api(
4387 access_token: str | None = None,
4388 ) -> list[str]:
4389 """
4390 Returns a list of `team_ids` the user belongs to from the Microsoft Graph API
4392 Args:
4393 access_token (Optional[str]): Microsoft Graph API access token
4395 Returns:
4396 List[str]: List of group IDs the user belongs to
4397 """
4398 try:
4399 async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER)
4401 # Handle MSFT Enterprise Application Groups
4402 service_principal_id: Final = os.getenv("MICROSOFT_SERVICE_PRINCIPAL_ID", None)
4403 service_principal_group_ids: list[str] | None = []
4404 service_principal_teams: list[MicrosoftServicePrincipalTeam] | None = []
4405 if service_principal_id:
4406 (
4407 service_principal_group_ids,
4408 service_principal_teams,
4409 ) = await MicrosoftSSOHandler.get_group_ids_from_service_principal(
4410 service_principal_id=service_principal_id,
4411 async_client=async_client,
4412 access_token=access_token,
4413 )
4414 verbose_proxy_logger.debug("Service principal group IDs: %s", service_principal_group_ids)
4415 if len(service_principal_group_ids) > 0:
4416 await MicrosoftSSOHandler.create_litellm_teams_from_service_principal_team_ids(
4417 service_principal_teams=service_principal_teams,
4418 )
4420 # Fetch user membership from Microsoft Graph API
4421 all_group_ids = []
4422 next_link: str | None = MicrosoftSSOHandler.get_graph_api_user_groups_endpoint()
4423 auth_headers: Final = {"Authorization": f"Bearer {access_token}"}
4424 page_count = 0
4426 while next_link is not None and page_count < MicrosoftSSOHandler.MAX_GRAPH_API_PAGES:
4427 group_ids, next_link = await MicrosoftSSOHandler.fetch_and_parse_groups(
4428 url=next_link, headers=auth_headers, async_client=async_client
4429 )
4430 all_group_ids.extend(group_ids)
4431 page_count += 1
4433 if next_link is not None and page_count >= MicrosoftSSOHandler.MAX_GRAPH_API_PAGES:
4434 verbose_proxy_logger.warning(
4435 "Reached maximum page limit of %s. Some groups may not be included.",
4436 MicrosoftSSOHandler.MAX_GRAPH_API_PAGES,
4437 )
4439 # If service_principal_group_ids is not empty, only return group_ids that are in both all_group_ids and service_principal_group_ids
4440 if service_principal_group_ids and len(service_principal_group_ids) > 0:
4441 all_group_ids = [group_id for group_id in all_group_ids if group_id in service_principal_group_ids]
4443 return all_group_ids
4445 except Exception as e:
4446 verbose_proxy_logger.error("Error getting user groups from Microsoft Graph API: %s", e)
4447 return []
4449 @staticmethod
4450 async def fetch_and_parse_groups(
4451 url: str, headers: dict, async_client: AsyncHTTPHandler
4452 ) -> tuple[list[str], str | None]:
4453 """Helper function to fetch and parse group data from a URL"""
4454 response: Final = await async_client.get(url, headers=headers)
4455 response_json: Final[dict[str, object]] = response.json()
4456 response_typed: Final = await MicrosoftSSOHandler._cast_graph_api_response_dict(response=response_json)
4457 group_ids: Final = MicrosoftSSOHandler._get_group_ids_from_graph_api_response(response=response_typed)
4458 return group_ids, response_typed.get("odata_nextLink")
4460 @staticmethod
4461 def _get_group_ids_from_graph_api_response(
4462 response: MicrosoftGraphAPIUserGroupResponse,
4463 ) -> list[str]:
4464 group_ids: Final = []
4465 for _object in response.get("value", []) or []:
4466 _group_id = _object.get("id")
4467 if _group_id is not None:
4468 group_ids.append(_group_id)
4469 return group_ids
4471 @staticmethod
4472 async def _cast_graph_api_response_dict(
4473 response: dict,
4474 ) -> MicrosoftGraphAPIUserGroupResponse:
4475 directory_objects: Final[list[MicrosoftGraphAPIUserGroupDirectoryObject]] = []
4476 for _object in response.get("value", []):
4477 directory_objects.append(
4478 MicrosoftGraphAPIUserGroupDirectoryObject(
4479 odata_type=_object.get("@odata.type"),
4480 id=_object.get("id"),
4481 deletedDateTime=_object.get("deletedDateTime"),
4482 description=_object.get("description"),
4483 displayName=_object.get("displayName"),
4484 roleTemplateId=_object.get("roleTemplateId"),
4485 )
4486 )
4487 return MicrosoftGraphAPIUserGroupResponse(
4488 odata_context=response.get("@odata.context"),
4489 odata_nextLink=response.get("@odata.nextLink"),
4490 value=directory_objects,
4491 )
4493 @staticmethod
4494 async def get_group_ids_from_service_principal(
4495 service_principal_id: str,
4496 async_client: AsyncHTTPHandler,
4497 access_token: str | None = None,
4498 ) -> tuple[list[str], list[MicrosoftServicePrincipalTeam]]:
4499 """
4500 Gets the groups belonging to the Service Principal Application
4502 Service Principal Id is an `Enterprise Application` in Azure AD
4504 Users use Enterprise Applications to manage Groups and Users on Microsoft Entra ID
4505 """
4506 base_url: Final = MicrosoftSSOHandler.get_graph_api_base_url()
4507 # Endpoint to get app role assignments for the given service principal
4508 endpoint: Final = f"/servicePrincipals/{service_principal_id}/appRoleAssignedTo"
4509 next_link: str | None = base_url + endpoint
4511 headers: Final = {
4512 "Authorization": f"Bearer {access_token}",
4513 "Content-Type": "application/json",
4514 }
4516 group_ids: Final[list[str]] = []
4517 service_principal_teams: Final[list[MicrosoftServicePrincipalTeam]] = []
4518 page_count = 0
4520 while next_link is not None and page_count < MicrosoftSSOHandler.MAX_GRAPH_API_PAGES:
4521 response = await async_client.get(next_link, headers=headers)
4522 response_json: _ServicePrincipalPage = response.json()
4523 verbose_proxy_logger.debug("Response from service principal app role assigned to: %s", response_json)
4525 for _object in response_json.get("value", []):
4526 if _object.get("principalType") == "Group":
4527 # Append the group ID to the list
4528 group_ids.append(_object.get("principalId"))
4529 # Append the service principal team to the list
4530 service_principal_teams.append(
4531 MicrosoftServicePrincipalTeam(
4532 principalDisplayName=_object.get("principalDisplayName"),
4533 principalId=_object.get("principalId"),
4534 )
4535 )
4537 next_link = response_json.get("@odata.nextLink")
4538 page_count += 1
4540 if next_link is not None and page_count >= MicrosoftSSOHandler.MAX_GRAPH_API_PAGES:
4541 verbose_proxy_logger.warning(
4542 "Reached maximum page limit of %s. Some service principal group assignments may not be included.",
4543 MicrosoftSSOHandler.MAX_GRAPH_API_PAGES,
4544 )
4546 return group_ids, service_principal_teams
4548 @staticmethod
4549 async def create_litellm_teams_from_service_principal_team_ids(
4550 service_principal_teams: list[MicrosoftServicePrincipalTeam],
4551 ):
4552 """
4553 Creates Litellm Teams from the Service Principal Group IDs
4555 When a user sets a `SERVICE_PRINCIPAL_ID` in the env, litellm will fetch groups under that service principal and create Litellm Teams from them
4556 """
4557 verbose_proxy_logger.debug("Creating Litellm Teams from Service Principal Teams: %s", service_principal_teams)
4558 for service_principal_team in service_principal_teams:
4559 litellm_team_id: str | None = service_principal_team.get("principalId")
4560 litellm_team_name: str | None = service_principal_team.get("principalDisplayName")
4561 if not litellm_team_id:
4562 verbose_proxy_logger.debug(
4563 "Skipping team creation for %s because it has no principalId", litellm_team_name
4564 )
4565 continue
4567 await SSOAuthenticationHandler.create_litellm_team_from_sso_group(
4568 litellm_team_id=litellm_team_id,
4569 litellm_team_name=litellm_team_name,
4570 )
4573class GoogleSSOHandler:
4574 """
4575 Handles Google SSO callback response and returns a CustomOpenID object
4576 """
4578 @staticmethod
4579 async def get_google_callback_response(
4580 request: Request,
4581 google_client_id: str,
4582 redirect_url: str,
4583 return_raw_sso_response: bool = False,
4584 ) -> OpenID | dict:
4585 """
4586 Get the Google SSO callback response
4588 Args:
4589 return_raw_sso_response: If True, return the raw SSO response
4590 """
4591 from fastapi_sso.sso.google import GoogleSSO
4593 google_client_secret: Final = os.getenv("GOOGLE_CLIENT_SECRET", None)
4594 if google_client_secret is None:
4595 raise ProxyException(
4596 message="GOOGLE_CLIENT_SECRET not set. Set it in .env file",
4597 type=ProxyErrorTypes.auth_error,
4598 param="GOOGLE_CLIENT_SECRET",
4599 code=status.HTTP_500_INTERNAL_SERVER_ERROR,
4600 )
4601 google_sso: Final = GoogleSSO(
4602 client_id=google_client_id,
4603 redirect_uri=redirect_url,
4604 client_secret=google_client_secret,
4605 )
4607 # if user is trying to get the raw sso response for debugging, return the raw sso response
4608 if return_raw_sso_response:
4609 return (
4610 await google_sso.verify_and_process(
4611 request=request,
4612 convert_response=False,
4613 )
4614 or {}
4615 )
4617 result: Final = await google_sso.verify_and_process(request)
4618 return result or {}
4621@router.get("/sso/debug/login", tags=["experimental"], include_in_schema=False)
4622async def debug_sso_login(request: Request):
4623 """
4624 Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env
4625 PROXY_BASE_URL should be the your deployed proxy endpoint, e.g. PROXY_BASE_URL="https://litellm-production-7002.up.railway.app/"
4626 Example:
4627 """
4628 from litellm.proxy.proxy_server import premium_user
4630 microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
4631 google_client_id: Final = os.getenv("GOOGLE_CLIENT_ID", None)
4632 generic_client_id: Final = os.getenv("GENERIC_CLIENT_ID", None)
4634 ####### Check if user is a Enterprise / Premium User #######
4635 if microsoft_client_id is not None or google_client_id is not None or generic_client_id is not None:
4636 if premium_user is not True:
4637 raise ProxyException(
4638 message="You must be a LiteLLM Enterprise user to use SSO. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this",
4639 type=ProxyErrorTypes.auth_error,
4640 param="premium_user",
4641 code=status.HTTP_403_FORBIDDEN,
4642 )
4644 # get url from request
4645 redirect_url: Final = SSOAuthenticationHandler.get_redirect_url_for_sso(
4646 request=request,
4647 sso_callback_route="sso/debug/callback",
4648 )
4650 # Check if we should use SSO handler
4651 if (
4652 SSOAuthenticationHandler.should_use_sso_handler(
4653 microsoft_client_id=microsoft_client_id,
4654 google_client_id=google_client_id,
4655 generic_client_id=generic_client_id,
4656 )
4657 is True
4658 ):
4659 return await SSOAuthenticationHandler.get_sso_login_redirect(
4660 redirect_url=redirect_url,
4661 microsoft_client_id=microsoft_client_id,
4662 google_client_id=google_client_id,
4663 generic_client_id=generic_client_id,
4664 request=request,
4665 )
4668@router.get("/sso/debug/callback", tags=["experimental"], include_in_schema=False)
4669async def debug_sso_callback(request: Request):
4670 """
4671 Returns the OpenID object returned by the SSO provider
4672 """
4673 import json
4675 from fastapi.responses import HTMLResponse
4677 from litellm.proxy._types import LiteLLM_JWTAuth
4678 from litellm.proxy.auth.handle_jwt import JWTHandler
4679 from litellm.proxy.proxy_server import (
4680 general_settings,
4681 jwt_handler,
4682 prisma_client,
4683 user_api_key_cache,
4684 )
4686 sso_jwt_handler: JWTHandler | None = None
4687 ui_access_mode: Final = general_settings.get("ui_access_mode", None)
4688 if ui_access_mode is not None and isinstance(ui_access_mode, dict):
4689 sso_jwt_handler = JWTHandler()
4690 sso_jwt_handler.update_environment(
4691 prisma_client=prisma_client,
4692 user_api_key_cache=user_api_key_cache,
4693 litellm_jwtauth=LiteLLM_JWTAuth(
4694 team_ids_jwt_field=general_settings.get("ui_access_mode", {}).get("sso_group_jwt_field", None),
4695 ),
4696 leeway=0,
4697 )
4699 microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
4700 google_client_id: Final = os.getenv("GOOGLE_CLIENT_ID", None)
4701 generic_client_id: Final = os.getenv("GENERIC_CLIENT_ID", None)
4703 redirect_url = os.getenv("PROXY_BASE_URL", str(request.base_url))
4704 if redirect_url.endswith("/"):
4705 redirect_url += "sso/debug/callback"
4706 else:
4707 redirect_url += "/sso/debug/callback"
4709 result = None
4710 received_response: dict | None = None
4711 access_token_payload: dict | None = None
4712 if google_client_id is not None:
4713 result = await GoogleSSOHandler.get_google_callback_response(
4714 request=request,
4715 google_client_id=google_client_id,
4716 redirect_url=redirect_url,
4717 return_raw_sso_response=True,
4718 )
4719 elif microsoft_client_id is not None:
4720 result = await MicrosoftSSOHandler.get_microsoft_callback_response(
4721 request=request,
4722 microsoft_client_id=microsoft_client_id,
4723 redirect_url=redirect_url,
4724 return_raw_sso_response=True,
4725 )
4727 elif generic_client_id is not None:
4728 (
4729 result,
4730 received_response,
4731 access_token_payload,
4732 _sso_assertion,
4733 ) = await get_generic_sso_response(
4734 request=request,
4735 jwt_handler=jwt_handler,
4736 generic_client_id=generic_client_id,
4737 redirect_url=redirect_url,
4738 sso_jwt_handler=sso_jwt_handler,
4739 )
4741 # If result is None, return a basic error message
4742 if result is None:
4743 return HTMLResponse(
4744 content="<h1>SSO Authentication Failed</h1><p>No data was returned from the SSO provider.</p>",
4745 status_code=400,
4746 )
4748 # Convert the OpenID object to a dictionary
4749 if hasattr(result, "__dict__"):
4750 result_dict = result.__dict__
4751 else:
4752 result_dict = dict(result)
4754 # Filter out any None values and convert to JSON serializable format
4755 filtered_result: Final = {}
4756 for key, value in result_dict.items():
4757 if value is not None and not key.startswith("_"):
4758 if isinstance(value, (str, int, float, bool)) or value is None:
4759 filtered_result[key] = value
4760 else:
4761 try:
4762 # Try to convert to string or another JSON serializable format
4763 filtered_result[key] = str(value)
4764 except Exception as e:
4765 filtered_result[key] = f"Complex value (not displayable): {e}"
4767 # Defense-in-depth: ensure no bearer tokens leak into the rendered HTML even if
4768 # a non-conforming IdP places them in its userinfo response.
4769 safe_raw_claims: Final = {k: v for k, v in (received_response or {}).items() if k not in _OAUTH_TOKEN_FIELDS}
4770 safe_access_token_claims = {k: v for k, v in (access_token_payload or {}).items() if k not in _OAUTH_TOKEN_FIELDS}
4772 await warn_if_id_jag_capture_gap()
4773 sso_payload: Final = {
4774 "parsed_by_proxy": filtered_result,
4775 "raw_claims": safe_raw_claims,
4776 "access_token_claims": safe_access_token_claims,
4777 }
4779 # Replace the placeholder in the template with the actual data
4780 sso_payload_json: Final = json.dumps(sso_payload, indent=2, default=str).replace("</", "<\\/")
4781 html_content: Final = jwt_display_template.replace(
4782 "const ssoData = SSO_DATA;",
4783 f"const ssoData = {sso_payload_json};",
4784 )
4786 return HTMLResponse(content=html_content)