fix(sso): direct PKCE token exchange + Redis wiring for multi-instance SSO (#22923)

* fix(sso): add direct PKCE token exchange and Redis cache wiring for multi-instance SSO

When PKCE is enabled, bypass fastapi-sso and perform direct token exchange so
code_verifier is correctly included. Store PKCE verifiers as dict in cache
for proper JSON serialization in Redis. Wire user_api_key_cache to Redis when
available so PKCE verifiers are shared across ECS tasks/pods.

Also adds clearer error messages when PKCE is required but not configured.

* refactor(sso): extract PKCE token exchange into SSOAuthenticationHandler methods

- Move import httpx/jwt to module level (top of file, not inside function)
- Extract inline PKCE token exchange + userinfo logic into two static methods:
  _pkce_token_exchange() and _get_pkce_userinfo()
- get_generic_sso_response PKCE path is now a single method call
- Fix double-logging in except block for non-PKCE errors
- Use %-style log formatting (no f-strings in log calls)

* fix: address greptile review feedback

- Fix access_token missing in PKCE path: read from combined_response directly
  instead of generic_sso.access_token (which is only set by verify_and_process)
- Fix PKCE error hint firing when PKCE is already enabled: only show
  'set GENERIC_CLIENT_USE_PKCE=true' advice when code_verifier was absent
- Fix unguarded KeyError on access_token: check for error field in HTTP 200
  responses before accessing token_response['access_token']
- Fix silent empty userinfo: raise ProxyException when both userinfo endpoint
  and id_token fallback produce no user data
- Fix backward-incompatible Redis wiring: only attach Redis to user_api_key_cache
  when GENERIC_CLIENT_USE_PKCE=true, preserving existing in-memory behaviour

* fix: address second round of greptile review feedback

- Fix PKCE error hint: check env var directly (not code_verifier presence) to
  distinguish 'PKCE not configured' from 'PKCE enabled but cache miss'
- Fix misleading Redis TTL comment in proxy_server.py

* fix: address third round of greptile review feedback

- Fix CRITICAL log firing on every non-PKCE callback: only log when PKCE is enabled
- Remove unused pkce_env_value intermediate variable
- Prefer reusing redis_usage_cache over creating separate RedisCache instance
  (avoids losing advanced connection options like SSL, timeouts, db)

* fix: address fourth round of greptile review feedback

- Strip OAuth token credentials from response_convertor input to prevent
  access_token/id_token appearing in restricted-group error messages
- Reuse single httpx.AsyncClient for both token exchange and userinfo requests
  to avoid a second TCP/TLS handshake per SSO callback
- Revert Redis wiring to user_api_key_cache: PKCE code already uses
  redis_usage_cache directly; wiring would route all API-key lookups through
  Redis unnecessarily. Add startup warning instead when PKCE+Redis mismatch.
- Move _OAUTH_TOKEN_FIELDS to module level

* fix remaining PKCE test assertion for dict-format verifier storage

* sanitize PKCE cache log to not expose verifier content

* address greptile review feedback (greploop iteration 3)

* address greptile review feedback (greploop iteration 4)

* address greptile review feedback (greploop iteration 5)

* simplify _get_pkce_userinfo: remove shared-client complexity, use async with directly

* address greptile review feedback (greploop iteration 6)

* address greptile review feedback (greploop iteration 7)

* address greptile review feedback (greploop iteration 8)

* address greptile review feedback (greploop iteration 9)

* address greptile review feedback (greploop iteration 10)

* address greptile review feedback (greploop iteration 11)

* address greptile review feedback (greploop iteration 12)

* fix misleading comment on user_api_key_cache TTL line

* address greptile review feedback (greploop iteration 13)

* address greptile review feedback (greploop iteration 14)

* address greptile review feedback (greploop iteration 15)

* address greptile review feedback (greploop iteration 16)

* address greptile review feedback (greploop iteration 17)

* address greptile review feedback (greploop iteration 18)

* address greptile review feedback (greploop iteration 19)

* address greptile review feedback (greploop iteration 20)

* address greptile review feedback (greploop iteration 21)

* address greptile review feedback (greploop iteration 22)

* address greptile review feedback (greploop iteration 23)

* address greptile review feedback (greploop iteration 24)

* address greptile review feedback (greploop iteration 25)

* address greptile review feedback (greploop iteration 26)

* address greptile review feedback (greploop iteration 27)

* address greptile review feedback (greploop iteration 28)

* address greptile review feedback (greploop iteration 29)

* address greptile review feedback (greploop iteration 30)

* address greptile review feedback (greploop iteration 31)

* address greptile review feedback (greploop iteration 32)

* address greptile review feedback (greploop iteration 33)

* address greptile review feedback (greploop iteration 34)

* address greptile review feedback (greploop iteration 35)

* address greptile review feedback (greploop iteration 37)

- read GENERIC_CLIENT_USE_PKCE env var once in prepare_token_exchange_parameters
- include actual decode error in jwt.decode failure exception message
- add GENERIC_CLIENT_USE_PKCE=true to no-state regression test

* defer PKCE verifier deletion until after all downstream processing

Move _delete_pkce_verifier to after response_convertor and
process_sso_jwt_access_token complete. If JWT processing raises,
the verifier stays in cache so the user can retry without restarting
the full OAuth flow.

* address greptile review feedback (greploop iteration 38)

- fix strict-mode cache miss error message to differentiate
  cross-instance routing failures (Redis configured) from single-instance
  issues (TTL expiry, pod restart) when only in-memory cache is available
- add comment above _get_pkce_userinfo call explaining that bearer
  credentials are always sourced from token_response in the merge step

* fix null JSON response body in _pkce_token_exchange

- Guard against HTTP 200 with body null: response.json() returns None
  for JSON null, and calling .get() on None raises AttributeError.
  Now raises a clean ProxyException with a clear error message.
- Fix misleading userinfo warning: was always saying "empty dict" but
  also fires for JSON null responses; updated to say "empty or null".
- Add HTTP status code assertion to cache miss test.

* address greptile review feedback (greploop iteration 39)

- fix credential leakage: directly assign received_response from
  combined_response instead of relying on nonlocal mutation; Pyright
  was flagging the old guard as unreachable, meaning credential stripping
  might not execute — now it always runs unconditionally
- add test for legacy plain-string cache format backward compat branch
- add test for HTTP 200 with no error field and no access_token (else branch)
- add test for HTTP 200 with JSON null body (new AttributeError guard)

* fix _OAUTH_TOKEN_FIELDS merge loop to preserve userinfo values on absent fields

When the token endpoint omits a bearer-credential field entirely (field
absent from token_response), the previous code deleted it from merged even
if userinfo provided a valid value. Now:
- non-null in token_response → restore authoritative token endpoint value
- explicit null in token_response → remove key from merged (clean absence)
- field absent from token_response → leave userinfo value unchanged

* use HTTP 401 for PKCE missing config errors

GENERIC_CLIENT_ID and GENERIC_TOKEN_ENDPOINT missing when PKCE is
enabled are auth-flow failures, not server errors. Use 401 instead
of 500 to avoid triggering false-positive server error alerts in
monitoring systems.

* address greptile review feedback (greploop iteration 40)

- fix duplicate error logging: demote first format-error log to DEBUG
  so the detailed ERROR in strict-mode branch is not duplicated
- add HTTP status code assertions to all PKCE ProxyException tests
  for better regression protection against accidental code changes

* add credential absence assertions to test_pkce_token_exchange_basic_auth

Verify that client_id and client_secret are NOT double-sent in the POST
body when Basic Auth is used (include_client_id=False with client_secret).
Catches regressions where credentials leak into both Auth header and body.

* address greptile review feedback (greploop iteration 41)

- add Bearer token header assertion to test_pkce_token_exchange_credentials_in_body
- add cache query assertions to both non-strict mode tests to confirm
  the cache was accessed before the warning path triggers

* address greptile review feedback (greploop iteration 42)

- assert null id_token is absent from merged result in basic auth test
- add test for HTTP 200 empty/null userinfo body with no id_token fallback

* use caplog to verify warning logs in non-strict cache miss tests

The two non-strict mode tests now use pytest's caplog fixture to assert
that a warning is actually emitted, not just that the code continues
without raising. This catches regressions where the warning silently
disappears.

* remove dead-code response=None guard in _pkce_token_exchange

* clean up stale pkce verifier cache entries in non-strict mode

* fix test: configure async_delete_cache as AsyncMock and assert cleanup called

* add sentinel guard so pkce-no-redis warning only fires once across hot-reloads

* fix misleading comments: code_verifier init and bearer-credential merge docs

* add best-effort cleanup in strict-mode for corrupt/empty cache entries

* add redirect_uri assertion, userinfo body in non-200 log, sentinel comment
This commit is contained in:
Ishaan Jaff 2026-03-12 12:41:39 -07:00 • committed by GitHub
parent 7ec78232b2
commit 7697b1c397
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 1355 additions and 79 deletions

View file

@ -17,6 +17,8 @@ import secrets
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast
import httpx
import jwt
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi.responses import RedirectResponse
@ -84,6 +86,7 @@ from litellm.proxy.utils import (
get_server_root_path,
)
from litellm.secret_managers.main import get_secret_bool, str_to_bool
from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403, F401
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
MicrosoftGraphAPIUserGroupDirectoryObject,
@ -92,7 +95,6 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
RoleMappings,
TeamMappings,
)
from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403, F401
from litellm.types.proxy.ui_sso import ParsedOpenIDResult
if TYPE_CHECKING:
@ -102,6 +104,12 @@ else:
router = APIRouter()
# OAuth bearer credential fields that must not appear in SSO debug responses
# (received_response is included in restricted-group error messages).
# Metadata fields (token_type, expires_in, scope) are intentionally kept so
# response convertors see the same fields in the PKCE path as in the non-PKCE path.
_OAUTH_TOKEN_FIELDS = frozenset({"access_token", "id_token", "refresh_token"})
def normalize_email(email: Optional[str]) -> Optional[str]:
"""
@ -774,25 +782,153 @@ async def get_generic_sso_response(
key, value = header.split("=")
additional_generic_sso_headers_dict[key] = value
code_verifier: Optional[str] = None # assigned inside try; initialized for type tracking
try:
result = await generic_sso.verify_and_process(
request,
params=await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=request,
generic_include_client_id=generic_include_client_id,
),
headers=additional_generic_sso_headers_dict,
token_exchange_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=request,
generic_include_client_id=generic_include_client_id,
)
access_token_str: Optional[str] = generic_sso.access_token
# Extract code_verifier (and the cache key for deferred deletion) before calling fastapi-sso
code_verifier = token_exchange_params.pop("code_verifier", None)
pkce_cache_key = token_exchange_params.pop("_pkce_cache_key", None)
# Get authorization code from query params (only used in the PKCE path below;
# the non-PKCE path delegates to verify_and_process which handles OAuth error
# callbacks — user-denied, CSRF mismatch — internally).
authorization_code = request.query_params.get("code")
if code_verifier:
if not authorization_code:
raise ProxyException(
message="Missing authorization code in callback",
type=ProxyErrorTypes.auth_error,
param="code",
code=status.HTTP_400_BAD_REQUEST,
)
if not generic_client_id:
raise ProxyException(
message="GENERIC_CLIENT_ID must be set when PKCE is enabled",
type=ProxyErrorTypes.auth_error,
param="GENERIC_CLIENT_ID",
code=status.HTTP_401_UNAUTHORIZED,
)
if not generic_token_endpoint:
raise ProxyException(
message="GENERIC_TOKEN_ENDPOINT must be set when PKCE is enabled",
type=ProxyErrorTypes.auth_error,
param="GENERIC_TOKEN_ENDPOINT",
code=status.HTTP_401_UNAUTHORIZED,
)
# All guards above raise, so authorization_code is a non-empty str here.
# Use an explicit type guard rather than assert (assert is a no-op with -O).
if not isinstance(authorization_code, str):
raise ProxyException(
message="Missing authorization code in callback",
type=ProxyErrorTypes.auth_error,
param="code",
code=status.HTTP_400_BAD_REQUEST,
)
combined_response = await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code=authorization_code,
code_verifier=code_verifier,
client_id=generic_client_id,
client_secret=generic_client_secret,
token_endpoint=generic_token_endpoint,
userinfo_endpoint=generic_userinfo_endpoint,
include_client_id=generic_include_client_id,
redirect_url=redirect_url,
additional_headers=additional_generic_sso_headers_dict,
)
# Pass the full response so custom response_convertor implementations
# can access all fields (including id_token for claim extraction).
result = response_convertor(combined_response, generic_sso)
# Strip bearer credentials from combined_response before storing in
# received_response. received_response may appear in restricted-group
# error messages — bearer tokens (access_token, id_token, refresh_token)
# must not be exposed to callers.
# Assign directly rather than relying on nonlocal mutation so that Pyright
# can track that received_response is non-None from this point on.
received_response = {
k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS
}
# In the PKCE path verify_and_process is skipped, so generic_sso.access_token
# is never set. Read the token directly from the exchange response instead so
# process_sso_jwt_access_token can extract JWT-embedded roles/teams.
access_token_str: Optional[str] = combined_response.get("access_token")
else:
result = await generic_sso.verify_and_process(
request,
params=token_exchange_params,
headers=additional_generic_sso_headers_dict,
)
access_token_str = generic_sso.access_token
process_sso_jwt_access_token(
access_token_str, sso_jwt_handler, result, role_mappings=role_mappings
)
# Delete the single-use PKCE verifier only after all downstream processing
# (response_convertor and process_sso_jwt_access_token) has completed
# successfully. Deleting earlier would consume the verifier on a transient
# failure, forcing the user to restart the entire OAuth flow from scratch.
if pkce_cache_key:
await SSOAuthenticationHandler._delete_pkce_verifier(pkce_cache_key)
except Exception as e:
verbose_proxy_logger.exception(
f"Error verifying and processing generic SSO: {e}. Passed in headers: {additional_generic_sso_headers_dict}"
)
error_message = str(e)
# Surface a helpful PKCE misconfiguration hint only when:
# 1. The error mentions PKCE/code verifier, AND
# 2. PKCE is not currently configured (GENERIC_CLIENT_USE_PKCE != true)
# If PKCE IS configured but code_verifier was absent (cross-instance cache miss),
# the real fix is shared Redis/sticky sessions — not enabling PKCE (it's already on).
pkce_configured = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
if not pkce_configured and (
"PKCE" in error_message or "code verifier" in error_message.lower()
):
is_okta = (
generic_authorization_endpoint
and "okta" in generic_authorization_endpoint.lower()
) or (generic_token_endpoint and "okta" in generic_token_endpoint.lower())
provider_name = "Okta" if is_okta else "Your OAuth provider"
detailed_message = (
f"SSO authentication failed: {provider_name} requires PKCE (Proof Key for Code Exchange) "
f"but it's not enabled in your LiteLLM configuration.\n\n"
f"SOLUTION: Add this environment variable and restart your proxy:\n"
f" GENERIC_CLIENT_USE_PKCE=true\n\n"
)
if is_okta:
detailed_message += (
"For AWS ECS: Add the environment variable to your task definition.\n"
"For Docker: Add -e GENERIC_CLIENT_USE_PKCE=true to your docker run command.\n"
"For .env file: Add GENERIC_CLIENT_USE_PKCE=true to your .env file.\n\n"
)
detailed_message += f"Original error: {error_message}"
raise ProxyException(
message=detailed_message,
type=ProxyErrorTypes.auth_error,
param="GENERIC_CLIENT_USE_PKCE",
code=status.HTTP_401_UNAUTHORIZED,
)
# Use .error() (not .exception()) for ProxyException — those are expected,
# intentional auth failures; emitting a full stack trace would produce
# false-positive alerts and pollute log aggregators.
if isinstance(e, ProxyException):
verbose_proxy_logger.error(
"SSO authentication failed: %s. Passed in headers: %s",
e,
additional_generic_sso_headers_dict,
)
else:
verbose_proxy_logger.exception(
"Error verifying and processing generic SSO: %s. Passed in headers: %s",
e,
additional_generic_sso_headers_dict,
)
raise e
verbose_proxy_logger.debug("generic result: %s", result)
return result or {}, received_response
@ -1773,21 +1909,25 @@ class SSOAuthenticationHandler:
# If PKCE is enabled, add PKCE parameters to the redirect URL
if code_verifier and "state" in redirect_params:
# Store code_verifier in cache (10 min TTL). Use Redis when available
# so callbacks landing on another pod can retrieve it (multi-pod SSO).
# Store code_verifier in cache (10 min TTL). Wrap in dict for proper
# JSON serialization in Redis. Use Redis when available so callbacks
# landing on another pod can retrieve it (multi-pod SSO).
cache_key = f"pkce_verifier:{redirect_params['state']}"
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(
key=cache_key,
value=code_verifier,
value={"code_verifier": code_verifier},
ttl=600,
)
else:
await user_api_key_cache.async_set_cache(
key=cache_key,
value=code_verifier,
value={"code_verifier": code_verifier},
ttl=600,
)
verbose_proxy_logger.debug(
"PKCE code_verifier stored in cache (TTL: 600s)"
)
# Add PKCE parameters to the authorization URL
if pkce_params:
@ -1813,9 +1953,6 @@ class SSOAuthenticationHandler:
# Update the redirect response
redirect_response.headers["location"] = new_url
verbose_proxy_logger.debug(
"PKCE parameters added to authorization URL"
)
return redirect_response
@staticmethod
@ -1859,6 +1996,7 @@ class SSOAuthenticationHandler:
# Handle PKCE (Proof Key for Code Exchange) if enabled
# Set GENERIC_CLIENT_USE_PKCE=true to enable PKCE for enhanced OAuth security
use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
if use_pkce:
(
code_verifier,
@ -1866,9 +2004,7 @@ class SSOAuthenticationHandler:
) = SSOAuthenticationHandler.generate_pkce_params()
redirect_params["code_challenge"] = code_challenge
redirect_params["code_challenge_method"] = "S256"
verbose_proxy_logger.debug(
"PKCE enabled - code_challenge added to authorization request"
)
verbose_proxy_logger.debug("PKCE enabled for authorization request")
return redirect_params, code_verifier
@ -2402,36 +2538,184 @@ class SSOAuthenticationHandler:
token_params: Dict[str, Any] = {"include_client_id": generic_include_client_id}
# Retrieve PKCE code_verifier if PKCE was used in authorization.
# Use same cache as store: Redis when available (multi-pod), else in-memory.
# Gate on GENERIC_CLIENT_USE_PKCE to avoid an unnecessary Redis round-trip
# on every non-PKCE SSO callback.
query_params = dict(request.query_params)
state = query_params.get("state")
if state:
use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
if use_pkce and not state:
verbose_proxy_logger.warning(
"PKCE is enabled (GENERIC_CLIENT_USE_PKCE=true) but no 'state' parameter "
"was found in the callback. The PKCE verifier cannot be retrieved without "
"a state value — the token exchange will proceed without code_verifier, "
"which the provider may reject. Ensure your OAuth provider returns 'state' "
"in the callback redirect."
)
if state and use_pkce:
from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
cache_key = f"pkce_verifier:{state}"
if redis_usage_cache is not None:
code_verifier = await redis_usage_cache.async_get_cache(key=cache_key)
cached_data = await redis_usage_cache.async_get_cache(key=cache_key)
else:
code_verifier = await user_api_key_cache.async_get_cache(key=cache_key)
cached_data = await user_api_key_cache.async_get_cache(key=cache_key)
code_verifier = None
# Track why code_verifier is absent for accurate strict-mode diagnostics.
_empty_value_in_dict = False # dict format correct but value is empty/null
if cached_data:
# Extract code_verifier from dict (stored as dict for JSON serialization)
if isinstance(cached_data, dict) and "code_verifier" in cached_data:
code_verifier = cached_data["code_verifier"]
if not code_verifier:
# Dict format is correct but value is empty or null. This is
# a distinct case from an unrecognized format — the entry exists
# but was stored with an empty/null verifier (data integrity issue).
_empty_value_in_dict = True
verbose_proxy_logger.warning(
"PKCE verifier dict for state '%s' has an empty/null code_verifier "
"value — may indicate a storage bug. Treating as a cache miss.",
state,
)
else:
verbose_proxy_logger.debug("PKCE code_verifier retrieved from cache")
elif isinstance(cached_data, str):
# Handle legacy format (plain string) for backward compatibility
code_verifier = cached_data
verbose_proxy_logger.warning(
"Retrieved code_verifier in legacy plain-string format. "
"Future storage will use dict format."
)
else:
# Defer the detailed ERROR log to the strict-mode branch below
# (which includes state and a diagnostic message). Log at DEBUG
# here to avoid duplicate ERROR entries in the same request.
verbose_proxy_logger.debug(
"Unexpected PKCE verifier cache format (type=%s); skipping.",
type(cached_data).__name__,
)
if code_verifier:
# Add code_verifier to token exchange parameters (Redis returns decoded string)
token_params["code_verifier"] = (
code_verifier
if isinstance(code_verifier, str)
else str(code_verifier)
# Add code_verifier to token exchange parameters.
token_params["code_verifier"] = code_verifier
# Return the cache key so the caller can delete it *after* a
# successful token exchange (avoids losing the verifier on retry
# if the exchange fails partway through).
token_params["_pkce_cache_key"] = cache_key
else:
# PKCE is enabled (already checked above) but verifier is missing.
# Most likely cause: callback landed on a different pod than the login
# request, and no shared Redis cache is configured.
active_cache = redis_usage_cache if redis_usage_cache is not None else user_api_key_cache
strict_cache_miss = (
os.getenv("PKCE_STRICT_CACHE_MISS", "false").lower() == "true"
)
verbose_proxy_logger.debug(
"PKCE code_verifier retrieved and will be included in token exchange"
)
# Clean up the cache entry (single-use verifier)
if redis_usage_cache is not None:
await redis_usage_cache.async_delete_cache(key=cache_key)
if strict_cache_miss:
# Distinguish empty-value dicts, corrupt-format entries, and genuine
# cache misses so operators can investigate the correct root cause.
if _empty_value_in_dict:
# Dict format was correct but code_verifier was empty/null.
# Best-effort cleanup: remove the corrupt entry before failing.
await SSOAuthenticationHandler._delete_pkce_verifier(cache_key)
raise ProxyException(
message=(
f"PKCE verifier for state '{state}' was found in cache but "
f"has an empty or null code_verifier value — possible storage bug."
),
type=ProxyErrorTypes.auth_error,
param="PKCE_CACHE_MISS",
code=status.HTTP_401_UNAUTHORIZED,
)
elif cached_data is not None:
# Cache had data but in an unrecognised format (e.g. corrupt Redis value).
# Best-effort cleanup: remove the corrupt entry before failing.
await SSOAuthenticationHandler._delete_pkce_verifier(cache_key)
verbose_proxy_logger.error(
"PKCE verifier for state '%s' has an unrecognized format (type=%s); "
"treating as a cache miss. Investigate the cached value — it may be "
"a corrupt or stale entry.",
state,
type(cached_data).__name__,
)
raise ProxyException(
message=(
f"PKCE verifier for state '{state}' has an unrecognized format "
f"(type={type(cached_data).__name__}). The cached entry may be corrupt."
),
type=ProxyErrorTypes.auth_error,
param="PKCE_CACHE_MISS",
code=status.HTTP_401_UNAUTHORIZED,
)
else:
# Genuine cache miss — verifier was never stored or already expired.
# Distinguish the likely cause: cross-instance routing (Redis configured
# but callback landed on a pod that never stored the verifier) vs.
# single-instance issues (TTL expiry, pod restart, or PKCE flow never
# started) when only in-memory cache is available.
if redis_usage_cache is not None:
cause = (
"The authorization and callback were likely handled by different "
"instances — the verifier was stored on one pod but not found on another."
)
else:
cause = (
"The verifier may have expired (TTL), been lost on a pod restart, "
"or the PKCE authorization step was never completed. "
"Configure Redis so all proxy instances share the PKCE verifier."
)
verbose_proxy_logger.error(
"PKCE is enabled but no verifier found in cache for state '%s'. "
"%s Cache type: %s.",
state,
cause,
type(active_cache).__name__,
)
raise ProxyException(
message=f"PKCE verifier not found in cache for state '{state}'. {cause}",
type=ProxyErrorTypes.auth_error,
param="PKCE_CACHE_MISS",
code=status.HTTP_401_UNAUTHORIZED,
)
else:
await user_api_key_cache.async_delete_cache(key=cache_key)
# Best-effort cleanup: if a stale/corrupt entry is present, delete it
# now so it does not linger until TTL expiry (resource hygiene).
if cached_data is not None:
await SSOAuthenticationHandler._delete_pkce_verifier(cache_key)
verbose_proxy_logger.warning(
"PKCE is enabled but verifier not found in cache for state '%s' "
"(cache type: %s, raw data present: %s). "
"Continuing without code_verifier — set PKCE_STRICT_CACHE_MISS=true to fail fast instead.",
state,
type(active_cache).__name__,
cached_data is not None,
)
return token_params
@staticmethod
async def _delete_pkce_verifier(cache_key: str) -> None:
"""Delete a single-use PKCE verifier from cache after a successful exchange.
Failure is non-fatal: a leftover verifier is a minor security concern
(unused key in cache) but not worth aborting an otherwise-successful login.
"""
from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache
try:
if redis_usage_cache is not None:
await redis_usage_cache.async_delete_cache(key=cache_key)
else:
await user_api_key_cache.async_delete_cache(key=cache_key)
except Exception as exc:
verbose_proxy_logger.warning(
"PKCE: failed to delete verifier cache key '%s' (best-effort cleanup): %s",
cache_key,
exc,
)
@staticmethod
def generate_pkce_params() -> Tuple[str, str]:
"""
@ -2460,6 +2744,303 @@ class SSOAuthenticationHandler:
return code_verifier, code_challenge
@staticmethod
async def _pkce_token_exchange(
authorization_code: str,
code_verifier: str,
client_id: str,
client_secret: Optional[str],
token_endpoint: str,
userinfo_endpoint: Optional[str],
include_client_id: bool,
redirect_url: Optional[str],
additional_headers: Dict[str, str],
) -> dict:
"""
Performs a direct OAuth token exchange including the PKCE code_verifier.
fastapi-sso does not forward code_verifier, so when PKCE is enabled we
bypass it and call the token endpoint ourselves, then fetch user info.
Returns a combined dict of the token response and user info, suitable
for passing to a response_convertor.
"""
verbose_proxy_logger.debug(
"PKCE: performing direct token exchange (code_verifier length=%d)",
len(code_verifier),
)
token_data: Dict[str, str] = {
"grant_type": "authorization_code",
"code": authorization_code,
"code_verifier": code_verifier,
}
# Only include redirect_uri when set — omitting it avoids sending the
# literal string "None" to the provider if the env var is missing.
if redirect_url:
token_data["redirect_uri"] = redirect_url
post_kwargs: Dict[str, Any] = {
"data": token_data,
"headers": {
**additional_headers,
"Content-Type": "application/x-www-form-urlencoded", # must not be overridden
"Accept": "application/json",
},
"timeout": 30.0,
}
if not include_client_id:
# Use Basic Auth only when a secret is available; public PKCE clients omit it.
if client_secret:
post_kwargs["auth"] = httpx.BasicAuth(client_id, client_secret)
else:
token_data["client_id"] = client_id
else:
token_data["client_id"] = client_id
if client_secret:
token_data["client_secret"] = client_secret
# The try/except is INSIDE the async with so that TLS teardown exceptions
# from __aexit__ propagate as-is and are NOT mis-labelled as "Token endpoint
# request failed". httpx buffers the full response body before __aexit__,
# so status_code / text / json() remain valid after the context exits.
async with httpx.AsyncClient() as http_client:
try:
response = await http_client.post(token_endpoint, **post_kwargs)
except Exception as exc:
# Catch network-level errors (SSL, DNS, TCP, timeout, etc.) and
# wrap them as a clean ProxyException rather than leaking raw
# httpx or OS exceptions to callers.
verbose_proxy_logger.error("PKCE token endpoint unreachable: %s", exc)
raise ProxyException(
message=f"Token endpoint request failed: {exc}",
type=ProxyErrorTypes.auth_error,
param="token_exchange",
code=status.HTTP_401_UNAUTHORIZED,
) from exc
# Response processing outside the async with — httpx buffers the full
# response body so status_code / text / json() remain valid after __aexit__.
if response.status_code != 200:
verbose_proxy_logger.error(
"PKCE token exchange failed. status=%s body=%s",
response.status_code,
response.text[:500],
)
raise ProxyException(
message=f"Token exchange failed: {response.status_code} - {response.text[:500]}",
type=ProxyErrorTypes.auth_error,
param="token_exchange",
code=status.HTTP_401_UNAUTHORIZED,
)
try:
token_response_raw = response.json()
except Exception as json_err:
verbose_proxy_logger.error(
"Failed to parse token response as JSON: %s. Body: %s",
json_err,
response.text[:500],
)
raise ProxyException(
message=f"Token endpoint returned invalid JSON: {json_err}",
type=ProxyErrorTypes.auth_error,
param="token_exchange",
code=status.HTTP_401_UNAUTHORIZED,
)
# Guard against HTTP 200 with body `null` — response.json() returns Python None
# in that case, and calling .get() on None raises AttributeError.
if not isinstance(token_response_raw, dict):
verbose_proxy_logger.error(
"Token endpoint returned non-dict JSON (type=%s). Body: %s",
type(token_response_raw).__name__,
response.text[:500],
)
raise ProxyException(
message=(
f"Token endpoint returned unexpected response format "
f"(expected JSON object, got {type(token_response_raw).__name__})"
),
type=ProxyErrorTypes.auth_error,
param="token_exchange",
code=status.HTTP_401_UNAUTHORIZED,
)
token_response: dict = token_response_raw
# Some providers return HTTP 200 with an error body (e.g. expired code, replay attack).
# Also guard against JSON `null` for access_token — it passes key-existence checks
# but would produce a "Bearer None" Authorization header downstream.
access_token_val = token_response.get("access_token")
if not isinstance(access_token_val, str) or not access_token_val:
error = token_response.get("error")
error_desc = token_response.get("error_description", "")
if error:
detail = f"{error} - {error_desc}" if error_desc else error
else:
detail = (
"token endpoint returned HTTP 200 but no access_token "
f"(response keys: {sorted(token_response.keys())})"
)
verbose_proxy_logger.error(
"Token response missing or null access_token. detail=%s", detail
)
raise ProxyException(
message=f"Token exchange failed: {detail}",
type=ProxyErrorTypes.auth_error,
param="token_exchange",
code=status.HTTP_401_UNAUTHORIZED,
)
verbose_proxy_logger.debug(
"PKCE token exchange successful. id_token_present=%s",
bool(token_response.get("id_token")),
)
# Bearer credentials (access_token, id_token, refresh_token) are always sourced
# from token_response — not from userinfo — in the merge step below.
userinfo = await SSOAuthenticationHandler._get_pkce_userinfo(
access_token=token_response["access_token"],
id_token=token_response.get("id_token"),
userinfo_endpoint=userinfo_endpoint,
additional_headers=additional_headers,
)
# Merge: userinfo takes precedence for identity claims (sub, email, name, …) per
# the OpenID Connect spec (userinfo is the authoritative source for identity).
# Bearer credentials (access_token, id_token, refresh_token) from the token endpoint
# take precedence over same-named fields in userinfo — non-standard providers sometimes
# include token fields in userinfo, which must not shadow the real bearer token.
# If a bearer field is absent from the token response, any userinfo-provided value
# is preserved as a fallback (useful for non-standard providers that omit id_token
# from the token response but include it in userinfo).
#
# Three-way merge semantics for each bearer-credential field:
# 1. token_response has a non-null value → use it (token endpoint is authoritative)
# 2. token_response explicitly sent null → remove the key so callers get a clean
# absence signal; the null from the token endpoint overrides userinfo too
# 3. field absent from token_response → leave whatever userinfo provided as-is
# (e.g. userinfo-provided id_token from a non-standard provider)
merged = {**token_response, **userinfo}
for field in _OAUTH_TOKEN_FIELDS:
if token_response.get(field) is not None:
# Case 1: non-null in token_response — restore authoritative value.
merged[field] = token_response[field]
elif field in token_response:
# Case 2: key exists but value is explicitly null — remove from merged.
merged.pop(field, None)
# Case 3: field absent from token_response — leave userinfo value as-is.
return merged
@staticmethod
async def _get_pkce_userinfo(
access_token: str,
id_token: Optional[str],
userinfo_endpoint: Optional[str],
additional_headers: Dict[str, str],
) -> dict:
"""
Fetches user info from the userinfo endpoint.
Falls back to decoding the id_token if the endpoint is unavailable.
"""
# None = request not yet attempted, failed, or returned empty/null (treated as failure
# so the id_token fallback can be attempted instead of returning a session with no claims).
userinfo: Optional[dict] = None
if userinfo_endpoint:
try:
async with httpx.AsyncClient() as client:
resp = await client.get(
userinfo_endpoint,
headers={
**additional_headers,
"Authorization": f"Bearer {access_token}", # must not be overridden
},
timeout=30.0,
)
if resp.status_code == 200:
try:
userinfo_raw = resp.json()
if not userinfo_raw:
# JSON null (None) or empty dict ({}) — no identity claims.
# Treat as failure so id_token fallback can be attempted.
verbose_proxy_logger.warning(
"Userinfo endpoint returned an empty or null response "
"(type=%s); treating as failure and attempting id_token fallback. "
"Check your provider's userinfo endpoint configuration.",
type(userinfo_raw).__name__,
)
userinfo = None
else:
userinfo = userinfo_raw
except Exception as json_err:
verbose_proxy_logger.warning(
"Userinfo endpoint returned non-JSON response (status 200): %s",
json_err,
)
else:
verbose_proxy_logger.warning(
"Userinfo endpoint returned %s (body: %s), falling back to id_token",
resp.status_code,
resp.text[:500],
)
except Exception as e:
verbose_proxy_logger.warning(
"Userinfo endpoint error: %s, falling back to id_token", e
)
# Only fall back to id_token when the userinfo request failed (None).
# Empty dict ({}) and JSON null are both treated as failure (set to None above) since
# they contain no identity claims — id_token fallback is attempted in that case too.
# Explicitly check for a non-empty string to avoid attempting JWT decode on
# a blank or non-string id_token field from a misbehaving provider.
if userinfo is None and isinstance(id_token, str) and id_token:
try:
userinfo = jwt.decode(id_token, options={"verify_signature": False})
if not userinfo:
# jwt.decode returned an empty dict (payload-free JWT or provider bug).
# Treat this the same as a missing userinfo — the session would have no
# identity claims, which is equivalent to a broken session.
verbose_proxy_logger.warning(
"id_token decoded to an empty payload — treating as failure."
)
userinfo = None
except Exception as decode_err:
verbose_proxy_logger.error("Failed to decode id_token: %s", decode_err)
raise ProxyException(
message=f"Failed to decode id_token JWT: {decode_err}",
type=ProxyErrorTypes.auth_error,
param="userinfo",
code=status.HTTP_401_UNAUTHORIZED,
)
if userinfo is None:
id_token_attempted = isinstance(id_token, str) and bool(id_token)
if userinfo_endpoint:
if id_token_attempted:
detail = (
"userinfo endpoint failed and id_token was present but "
"decoded to an empty payload — no identity claims available"
)
else:
detail = "userinfo endpoint failed and no id_token was present in the token response"
else:
if id_token_attempted:
detail = (
"no userinfo endpoint is configured (GENERIC_USERINFO_ENDPOINT) "
"and id_token decoded to an empty payload — no identity claims available"
)
else:
detail = "no userinfo endpoint is configured (GENERIC_USERINFO_ENDPOINT) and no id_token was present"
raise ProxyException(
message=f"SSO user info unavailable: {detail}.",
type=ProxyErrorTypes.auth_error,
param="userinfo",
code=status.HTTP_401_UNAUTHORIZED,
)
return userinfo
class MicrosoftSSOHandler:
"""

View file

@ -1541,6 +1541,9 @@ native_background_mode: List[
polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache
user_custom_auth = None
user_custom_key_generate = None
# Sentinel: prevents PKCE-no-Redis advisory from re-logging on config hot-reload.
# Tests that need to reset it can patch 'litellm.proxy.proxy_server._pkce_no_redis_warning_emitted'.
_pkce_no_redis_warning_emitted: bool = False
user_custom_sso = None
user_custom_ui_sso_sign_in_handler = None
use_background_health_checks = None
@ -2520,6 +2523,8 @@ class ProxyConfig:
):
## INIT PROXY REDIS USAGE CLIENT ##
redis_usage_cache = litellm.cache.cache
# Note: PKCE verifier storage uses redis_usage_cache directly (not
# user_api_key_cache) to avoid routing all API-key lookups through Redis.
def switch_on_llm_response_caching(self):
"""
@ -3066,8 +3071,28 @@ class ProxyConfig:
if user_api_key_cache_ttl is not None:
user_api_key_cache.update_cache_ttl(
default_in_memory_ttl=float(user_api_key_cache_ttl),
default_redis_ttl=None, # user_api_key_cache is an in-memory cache
default_redis_ttl=None, # user_api_key_cache uses in-memory TTL only; Redis not configured for key lookups
)
### PKCE MULTI-INSTANCE PREREQUISITE CHECK ###
# PKCE verifiers are stored in redis_usage_cache when available so they can
# be read back by any instance (not just the one that started the auth flow).
# user_api_key_cache is intentionally left in-memory-only to avoid routing
# all API-key lookups through Redis.
use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
if use_pkce and redis_usage_cache is None:
global _pkce_no_redis_warning_emitted
if not _pkce_no_redis_warning_emitted:
_pkce_no_redis_warning_emitted = True
verbose_proxy_logger.warning(
"GENERIC_CLIENT_USE_PKCE=true but Redis is not configured for LiteLLM caching. "
"PKCE verifiers will not be shared across instances — callbacks may land on a "
"different pod than the login request and fail silently. "
"Configure Redis via the 'cache' section in your proxy config, "
"or enable sticky sessions for single-instance deployments. "
"Set PKCE_STRICT_CACHE_MISS=true to fail fast with a 401 on cache misses "
"instead of continuing without a code_verifier."
)
### STORE MODEL IN DB ### feature flag for `/model/new`
store_model_in_db = general_settings.get("store_model_in_db", False)
if store_model_in_db is None:

View file

@ -4,6 +4,7 @@ import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import Request
@ -22,10 +23,10 @@ from litellm.proxy.management_endpoints.ui_sso import (
GoogleSSOHandler,
MicrosoftSSOHandler,
SSOAuthenticationHandler,
_setup_team_mappings,
determine_role_from_groups,
normalize_email,
process_sso_jwt_access_token,
determine_role_from_groups,
_setup_team_mappings,
)
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
@ -1791,8 +1792,8 @@ class TestCustomUISSO:
async def mock_google_login():
# This mimics the relevant part of google_login that would trigger the import error
try:
from enterprise.litellm_enterprise.proxy.auth.custom_sso_handler import (
EnterpriseCustomSSOHandler, # noqa: F401
from enterprise.litellm_enterprise.proxy.auth.custom_sso_handler import ( # noqa: F401
EnterpriseCustomSSOHandler,
)
return "success"
@ -3141,13 +3142,15 @@ class TestPKCEFunctionality:
test_state = "test_oauth_state_123"
mock_request.query_params = {"state": test_state}
# Mock cache with async methods
# Mock cache with async methods — use dict format (primary path)
mock_cache = MagicMock()
test_code_verifier = "test_code_verifier_abc123xyz"
mock_cache.async_get_cache = AsyncMock(return_value=test_code_verifier)
mock_cache.async_get_cache = AsyncMock(
return_value={"code_verifier": test_code_verifier}
)
mock_cache.async_delete_cache = AsyncMock()
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
# Act
token_params = (
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
@ -3158,14 +3161,15 @@ class TestPKCEFunctionality:
# Assert
assert token_params["include_client_id"] is False
assert token_params["code_verifier"] == test_code_verifier
# Cache key is returned for deferred deletion (after exchange succeeds)
assert token_params["_pkce_cache_key"] == f"pkce_verifier:{test_state}"
# Verify cache was accessed and deleted
# Verify cache was read but NOT deleted yet (deletion is deferred to after
# successful token exchange to preserve the verifier for retries)
mock_cache.async_get_cache.assert_called_once_with(
key=f"pkce_verifier:{test_state}"
)
mock_cache.async_delete_cache.assert_called_once_with(
key=f"pkce_verifier:{test_state}"
)
mock_cache.async_delete_cache.assert_not_called()
@pytest.mark.asyncio
async def test_get_generic_sso_redirect_response_with_pkce(self):
@ -3191,7 +3195,9 @@ class TestPKCEFunctionality:
mock_cache.async_set_cache = AsyncMock()
with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}):
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
):
# Act
result = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_sso,
@ -3205,7 +3211,10 @@ class TestPKCEFunctionality:
cache_call = mock_cache.async_set_cache.call_args
assert cache_call.kwargs["key"] == f"pkce_verifier:{test_state}"
assert cache_call.kwargs["ttl"] == 600
assert len(cache_call.kwargs["value"]) == 43
# Value is stored as dict for proper JSON serialization in Redis
cached = cache_call.kwargs["value"]
assert isinstance(cached, dict) and "code_verifier" in cached
assert len(cached["code_verifier"]) == 43
# Verify PKCE parameters were added to the redirect URL
assert result is not None
@ -3273,7 +3282,10 @@ class TestPKCEFunctionality:
stored_key = "pkce_verifier:multi_pod_state_xyz"
assert stored_key in mock_redis._store
stored_value = mock_redis._store[stored_key]
assert isinstance(stored_value, str) and len(json.loads(stored_value)) == 43
# Stored as JSON-serialized dict for Redis compatibility
stored_dict = json.loads(stored_value)
assert isinstance(stored_dict, dict) and "code_verifier" in stored_dict
assert len(stored_dict["code_verifier"]) == 43
# Pod B: callback with same state, retrieve from "Redis"
mock_request = MagicMock(spec=Request)
@ -3282,12 +3294,12 @@ class TestPKCEFunctionality:
request=mock_request, generic_include_client_id=False
)
assert "code_verifier" in token_params
assert token_params["code_verifier"] == json.loads(stored_value)
assert token_params["code_verifier"] == stored_dict["code_verifier"]
# Cache key returned for deferred deletion after successful exchange
assert token_params["_pkce_cache_key"] == stored_key
mock_in_memory.async_get_cache.assert_not_called()
# delete_cache called; key removed (asserted below)
# Verifier consumed (single-use); key removed from "Redis"
assert "pkce_verifier:multi_pod_state_xyz" not in mock_redis._store
# Deletion is deferred — key still present until exchange succeeds
assert stored_key in mock_redis._store
@pytest.mark.asyncio
async def test_pkce_fallback_in_memory_roundtrip_when_redis_none(self):
@ -3342,7 +3354,7 @@ class TestPKCEFunctionality:
"value"
]
assert stored_key == "pkce_verifier:fallback_state_xyz"
assert isinstance(stored_value, str) and len(stored_value) == 43
assert isinstance(stored_value, dict) and len(stored_value["code_verifier"]) == 43
# Same pod: callback retrieves from in-memory cache
mock_request = MagicMock(spec=Request)
@ -3351,16 +3363,14 @@ class TestPKCEFunctionality:
request=mock_request, generic_include_client_id=False
)
assert "code_verifier" in token_params
assert token_params["code_verifier"] == stored_value
assert token_params["code_verifier"] == stored_value["code_verifier"]
# Cache key returned for deferred deletion after successful exchange
assert token_params["_pkce_cache_key"] == stored_key
mock_in_memory.async_get_cache.assert_called_once_with(
key=stored_key
)
mock_in_memory.async_delete_cache.assert_called_once_with(
key=stored_key
)
# Verifier consumed; key removed from in-memory
assert "pkce_verifier:fallback_state_xyz" not in in_memory_store
# Deletion is deferred — not called by prepare_token_exchange_parameters
mock_in_memory.async_delete_cache.assert_not_called()
@pytest.mark.asyncio
async def test_pkce_prepare_token_exchange_returns_nothing_when_no_state(self):
@ -3373,18 +3383,677 @@ class TestPKCEFunctionality:
mock_redis = MagicMock()
mock_in_memory = MagicMock()
with patch("litellm.proxy.proxy_server.redis_usage_cache", mock_redis):
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory):
mock_request = MagicMock(spec=Request)
mock_request.query_params = {}
token_params = (
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
with patch("litellm.proxy.proxy_server.redis_usage_cache", mock_redis), patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_in_memory
), patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}, clear=False):
mock_request = MagicMock(spec=Request)
mock_request.query_params = {}
token_params = (
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert "code_verifier" not in token_params
mock_redis.async_get_cache.assert_not_called()
mock_in_memory.async_get_cache.assert_not_called()
)
assert "code_verifier" not in token_params
mock_redis.async_get_cache.assert_not_called()
mock_in_memory.async_get_cache.assert_not_called()
@pytest.mark.asyncio
async def test_pkce_token_exchange_basic_auth(self):
"""When include_client_id=False, client credentials go via HTTP Basic Auth."""
token_resp = {
"access_token": "tok_abc",
"id_token": None,
"token_type": "Bearer",
"expires_in": 3600,
}
userinfo_resp = {"sub": "user1", "email": "user@example.com"}
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = token_resp
mock_userinfo_response = MagicMock()
mock_userinfo_response.status_code = 200
mock_userinfo_response.json.return_value = userinfo_resp
async def fake_post(*args, **kwargs):
# Verify Basic Auth is set
assert "auth" in kwargs
assert isinstance(kwargs["auth"], httpx.BasicAuth)
# Verify code_verifier is in the POST body (essential PKCE field)
post_data = kwargs.get("data", {})
assert post_data.get("code_verifier") == "verifier_abc"
# Verify redirect_uri is forwarded (required by strict OAuth providers)
assert post_data.get("redirect_uri") == "https://proxy.example.com/callback"
# Verify credentials are NOT double-sent in the POST body when using Basic Auth
assert "client_secret" not in post_data, "client_secret must not appear in POST body when using Basic Auth"
assert "client_id" not in post_data, "client_id must not appear in POST body when using Basic Auth (include_client_id=False)"
return mock_response
# Use separate mock clients for token exchange and userinfo —
# each httpx.AsyncClient() call gets its own independent mock.
mock_token_client = AsyncMock()
mock_token_client.__aenter__ = AsyncMock(return_value=mock_token_client)
mock_token_client.__aexit__ = AsyncMock(return_value=False)
mock_token_client.post = AsyncMock(side_effect=fake_post)
mock_userinfo_client = AsyncMock()
mock_userinfo_client.__aenter__ = AsyncMock(return_value=mock_userinfo_client)
mock_userinfo_client.__aexit__ = AsyncMock(return_value=False)
mock_userinfo_client.get = AsyncMock(return_value=mock_userinfo_response)
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client_cls.side_effect = [mock_token_client, mock_userinfo_client]
result = await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="auth_code_123",
code_verifier="verifier_abc",
client_id="my_client",
client_secret="my_secret",
token_endpoint="https://example.com/token",
userinfo_endpoint="https://example.com/userinfo",
include_client_id=False,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert result["access_token"] == "tok_abc"
assert result["email"] == "user@example.com"
# id_token was explicit null in token_response — the merge loop must remove it
# rather than leaving "id_token": None in the result.
assert "id_token" not in result, "null id_token from token endpoint must be absent in merged result"
# Verify userinfo GET used the correct Bearer token header
get_call = mock_userinfo_client.get.call_args
assert get_call is not None
assert get_call.kwargs["headers"]["Authorization"] == "Bearer tok_abc"
@pytest.mark.asyncio
async def test_pkce_token_exchange_credentials_in_body(self):
"""When include_client_id=True, credentials go in the request body."""
token_resp = {
"access_token": "tok_body",
"id_token": None,
"token_type": "Bearer",
"expires_in": 3600,
}
userinfo_resp = {"sub": "user2", "email": "user2@example.com"}
async def fake_post(*args, **kwargs):
assert "auth" not in kwargs, "Should NOT use Basic Auth when include_client_id=True"
data = kwargs.get("data", {})
assert "client_id" in data
assert "client_secret" in data
assert data.get("code_verifier") == "verifier_xyz", "code_verifier must be in POST body"
assert data.get("redirect_uri") == "https://proxy.example.com/callback", "redirect_uri must be forwarded"
mock = MagicMock()
mock.status_code = 200
mock.json.return_value = token_resp
return mock
mock_userinfo = MagicMock()
mock_userinfo.status_code = 200
mock_userinfo.json.return_value = userinfo_resp
mock_token_client = AsyncMock()
mock_token_client.__aenter__ = AsyncMock(return_value=mock_token_client)
mock_token_client.__aexit__ = AsyncMock(return_value=False)
mock_token_client.post = AsyncMock(side_effect=fake_post)
mock_userinfo_client = AsyncMock()
mock_userinfo_client.__aenter__ = AsyncMock(return_value=mock_userinfo_client)
mock_userinfo_client.__aexit__ = AsyncMock(return_value=False)
mock_userinfo_client.get = AsyncMock(return_value=mock_userinfo)
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client_cls.side_effect = [mock_token_client, mock_userinfo_client]
result = await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="auth_code_456",
code_verifier="verifier_xyz",
client_id="client_id_value",
client_secret="client_secret_value",
token_endpoint="https://example.com/token",
userinfo_endpoint="https://example.com/userinfo",
include_client_id=True,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert result["access_token"] == "tok_body"
assert result["sub"] == "user2"
# Verify userinfo GET used the correct Bearer token header
get_call = mock_userinfo_client.get.call_args
assert get_call is not None
assert get_call.kwargs["headers"]["Authorization"] == "Bearer tok_body"
@pytest.mark.asyncio
async def test_pkce_token_exchange_http200_with_error_body(self):
"""Provider returns HTTP 200 but with an error field instead of tokens."""
from litellm.proxy._types import ProxyException
error_body = {"error": "invalid_grant", "error_description": "Code already used"}
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.json.return_value = error_body
mock_client.post = AsyncMock(return_value=mock_resp)
mock_client_cls.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="expired_code",
code_verifier="verifier",
client_id="cid",
client_secret="csecret",
token_endpoint="https://example.com/token",
userinfo_endpoint="https://example.com/userinfo",
include_client_id=False,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert "invalid_grant" in exc_info.value.message
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_userinfo_falls_back_to_id_token(self):
"""When the userinfo endpoint fails, decode the id_token as fallback."""
import base64
import json as _json
payload = {"sub": "user_from_jwt", "email": "jwt@example.com"}
# Build a minimal JWT (header.payload.signature — signature not verified)
encoded_payload = base64.urlsafe_b64encode(
_json.dumps(payload).encode()
).rstrip(b"=").decode()
fake_id_token = f"eyJhbGciOiJSUzI1NiJ9.{encoded_payload}.fakesig"
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
mock_fail = MagicMock()
mock_fail.status_code = 503
mock_client.get = AsyncMock(return_value=mock_fail)
mock_client_cls.return_value = mock_client
result = await SSOAuthenticationHandler._get_pkce_userinfo(
access_token="some_token",
id_token=fake_id_token,
userinfo_endpoint="https://example.com/userinfo",
additional_headers={},
)
assert result["sub"] == "user_from_jwt"
assert result["email"] == "jwt@example.com"
@pytest.mark.asyncio
async def test_pkce_userinfo_uses_id_token_when_no_endpoint(self):
"""When userinfo_endpoint is None, fall back to id_token directly without HTTP call."""
import base64
import json as _json
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
payload = {"sub": "id_token_user", "email": "id@example.com"}
encoded_payload = (
base64.urlsafe_b64encode(_json.dumps(payload).encode()).rstrip(b"=").decode()
)
fake_id_token = f"eyJhbGciOiJSUzI1NiJ9.{encoded_payload}.fakesig"
# No httpx call should happen when userinfo_endpoint is None
result = await SSOAuthenticationHandler._get_pkce_userinfo(
access_token="some_token",
id_token=fake_id_token,
userinfo_endpoint=None,
additional_headers={},
)
assert result["sub"] == "id_token_user"
assert result["email"] == "id@example.com"
@pytest.mark.asyncio
async def test_pkce_userinfo_raises_when_both_sources_unavailable(self):
"""When userinfo endpoint fails AND no id_token, raise ProxyException."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
mock_fail = MagicMock()
mock_fail.status_code = 503
mock_client.get = AsyncMock(return_value=mock_fail)
mock_client_cls.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._get_pkce_userinfo(
access_token="token",
id_token=None, # no id_token available
userinfo_endpoint="https://example.com/userinfo",
additional_headers={},
)
assert "unavailable" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_userinfo_http200_empty_body_no_id_token_raises(self):
"""When userinfo returns HTTP 200 with an empty/null body and no id_token is
available, _get_pkce_userinfo raises ProxyException."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.json.return_value = None # HTTP 200 with null JSON body
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
mock_client.get = AsyncMock(return_value=mock_resp)
mock_client_cls.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._get_pkce_userinfo(
access_token="access_token",
id_token=None, # no id_token fallback available
userinfo_endpoint="https://example.com/userinfo",
additional_headers={},
)
assert "unavailable" in exc_info.value.message.lower() or "no userinfo" in exc_info.value.message.lower() or "userinfo" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_cache_miss_raises_proxy_exception(self):
"""prepare_token_exchange_parameters raises ProxyException when PKCE is enabled
but no verifier is found in cache (cross-instance cache miss scenario)."""
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "missing_state_123"}
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
), patch.dict(
os.environ,
{"GENERIC_CLIENT_USE_PKCE": "true", "PKCE_STRICT_CACHE_MISS": "true"},
):
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert "verifier not found" in exc_info.value.message.lower() or "cache" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_token_exchange_public_client_no_secret(self):
"""Public PKCE client (include_client_id=False, no secret) sends client_id in
POST body and does NOT include Basic Auth or client_secret."""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
token_resp = {
"access_token": "tok_public",
"token_type": "Bearer",
"expires_in": 3600,
}
userinfo_resp = {"sub": "pubuser", "email": "pub@example.com"}
async def fake_post(*args, **kwargs):
assert "auth" not in kwargs, "Public client must not use Basic Auth"
data = kwargs.get("data", {})
assert data.get("client_id") == "public_client_id"
assert "client_secret" not in data, "No secret should be sent for public client"
assert data.get("code_verifier") == "public_verifier"
mock = MagicMock()
mock.status_code = 200
mock.json.return_value = token_resp
return mock
mock_userinfo = MagicMock()
mock_userinfo.status_code = 200
mock_userinfo.json.return_value = userinfo_resp
mock_token_client = AsyncMock()
mock_token_client.__aenter__ = AsyncMock(return_value=mock_token_client)
mock_token_client.__aexit__ = AsyncMock(return_value=False)
mock_token_client.post = AsyncMock(side_effect=fake_post)
mock_userinfo_client = AsyncMock()
mock_userinfo_client.__aenter__ = AsyncMock(return_value=mock_userinfo_client)
mock_userinfo_client.__aexit__ = AsyncMock(return_value=False)
mock_userinfo_client.get = AsyncMock(return_value=mock_userinfo)
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client_cls.side_effect = [mock_token_client, mock_userinfo_client]
result = await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="auth_pub",
code_verifier="public_verifier",
client_id="public_client_id",
client_secret=None, # public client — no secret
token_endpoint="https://example.com/token",
userinfo_endpoint="https://example.com/userinfo",
include_client_id=False,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert result["access_token"] == "tok_public"
assert result["sub"] == "pubuser"
@pytest.mark.asyncio
async def test_delete_pkce_verifier_swallows_deletion_errors(self):
"""_delete_pkce_verifier must not raise when the cache delete fails
(best-effort cleanup — a leftover verifier must not abort a successful SSO login)."""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
failing_cache = MagicMock()
failing_cache.async_delete_cache = AsyncMock(side_effect=Exception("Redis down"))
# Should NOT raise even though the underlying cache delete fails
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch(
"litellm.proxy.proxy_server.user_api_key_cache", failing_cache
):
await SSOAuthenticationHandler._delete_pkce_verifier("pkce_verifier:test_state")
failing_cache.async_delete_cache.assert_called_once_with(key="pkce_verifier:test_state")
@pytest.mark.asyncio
async def test_pkce_cache_miss_unexpected_format_raises_proxy_exception(self):
"""When cached data exists but has an unrecognized format (not a dict with
code_verifier, not a plain string), prepare_token_exchange_parameters raises
ProxyException rather than silently falling through to a non-PKCE flow."""
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Cache returns an integer — unexpected format
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=12345)
mock_cache.async_delete_cache = AsyncMock()
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "bad_format_state"}
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
), patch.dict(
os.environ,
{"GENERIC_CLIENT_USE_PKCE": "true", "PKCE_STRICT_CACHE_MISS": "true"},
):
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert "cache" in exc_info.value.message.lower() or "verifier" in exc_info.value.message.lower() or "format" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
# Strict mode should also clean up the corrupt cache entry before raising
mock_cache.async_delete_cache.assert_called_once()
@pytest.mark.asyncio
async def test_pkce_cache_miss_non_strict_logs_warning_and_continues(self, caplog):
"""Default (non-strict) cache-miss behavior: logs a warning and returns params
without code_verifier rather than raising, to preserve backward compatibility."""
import logging
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "missing_state_non_strict"}
# PKCE_STRICT_CACHE_MISS explicitly set to false — should NOT raise.
# Use patch.dict with the key set to "false" rather than os.environ.pop()
# to avoid permanently mutating the test process environment.
with caplog.at_level(logging.WARNING), patch(
"litellm.proxy.proxy_server.redis_usage_cache", None
), patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
), patch.dict(
os.environ,
{"GENERIC_CLIENT_USE_PKCE": "true", "PKCE_STRICT_CACHE_MISS": "false"},
clear=False,
):
result = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
# Should return params without code_verifier (no raise)
assert "code_verifier" not in result
assert "_pkce_cache_key" not in result
# Non-strict mode emits a warning rather than raising
mock_cache.async_get_cache.assert_called_once()
# Verify the warning was actually logged
assert any(
"verifier not found" in r.message.lower() or "code_verifier" in r.message.lower()
for r in caplog.records
if r.levelno >= logging.WARNING
), f"Expected a cache-miss warning. Records: {[r.message for r in caplog.records]}"
@pytest.mark.asyncio
async def test_pkce_token_exchange_non200_raises_proxy_exception(self):
"""_pkce_token_exchange raises ProxyException when the token endpoint
returns a non-200 status (e.g. 401 Unauthorized from provider)."""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.text = "Unauthorized"
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
mock_client.post = AsyncMock(return_value=mock_response)
mock_client_cls.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="auth_code",
code_verifier="verifier",
client_id="client_id",
client_secret="secret",
token_endpoint="https://example.com/token",
userinfo_endpoint=None,
include_client_id=True,
redirect_url="https://proxy.example.com/callback",
additional_headers={},
)
assert "token" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_cache_miss_unexpected_format_non_strict_logs_warning(self, caplog):
"""When cached data has an unexpected format (e.g. integer from corrupt Redis)
in non-strict mode, prepare_token_exchange_parameters logs a warning and
returns params without code_verifier rather than raising."""
import logging
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Cache returns an integer — unexpected format
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=12345)
mock_cache.async_delete_cache = AsyncMock()
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "bad_format_non_strict"}
# Non-strict mode: should log a warning and continue, not raise.
# Use patch.dict with PKCE_STRICT_CACHE_MISS="false" to avoid permanently
# mutating the test process environment with os.environ.pop().
with caplog.at_level(logging.WARNING), patch(
"litellm.proxy.proxy_server.redis_usage_cache", None
), patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
), patch.dict(
os.environ,
{"GENERIC_CLIENT_USE_PKCE": "true", "PKCE_STRICT_CACHE_MISS": "false"},
clear=False,
):
result = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
# No raise in non-strict mode; verifier simply absent from params
assert "code_verifier" not in result
assert "_pkce_cache_key" not in result
# Cache was queried (the unexpected format was retrieved and logged at WARNING)
mock_cache.async_get_cache.assert_called_once()
# Verify a warning was logged about the unexpected format or cache miss
assert any(
"verifier" in r.message.lower() or "format" in r.message.lower() or "cache" in r.message.lower()
for r in caplog.records
if r.levelno >= logging.WARNING
), f"Expected a format/cache warning. Records: {[r.message for r in caplog.records]}"
# Verify cleanup was attempted for the corrupt/stale cache entry
mock_cache.async_delete_cache.assert_called_once()
@pytest.mark.asyncio
async def test_pkce_legacy_string_cache_format_backward_compat(self):
"""Legacy plain-string cache entries (stored before dict format was introduced)
are handled transparently via the backward-compat branch."""
import os
from unittest.mock import AsyncMock, MagicMock, patch
from starlette.requests import Request
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
legacy_verifier = "legacy_plain_string_verifier_abc123"
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value=legacy_verifier)
mock_request = MagicMock(spec=Request)
mock_request.query_params = {"state": "legacy_state_xyz"}
with patch("litellm.proxy.proxy_server.redis_usage_cache", None), patch(
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
), patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}, clear=False):
result = await SSOAuthenticationHandler.prepare_token_exchange_parameters(
request=mock_request, generic_include_client_id=False
)
assert result["code_verifier"] == legacy_verifier
assert result["_pkce_cache_key"] == "pkce_verifier:legacy_state_xyz"
@pytest.mark.asyncio
async def test_pkce_token_exchange_null_json_body_raises_proxy_exception(self):
"""HTTP 200 with JSON body `null` raises a clean ProxyException instead of
AttributeError when .get() is called on the None return value."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.json.return_value = None # JSON null response body
mock_resp.text = "null"
mock_client.post = AsyncMock(return_value=mock_resp)
mock_client_cls.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="some_code",
code_verifier="verifier",
client_id="cid",
client_secret="csecret",
token_endpoint="https://example.com/token",
userinfo_endpoint=None,
include_client_id=False,
redirect_url=None,
additional_headers={},
)
assert "unexpected response format" in exc_info.value.message.lower()
assert str(exc_info.value.code) == "401"
@pytest.mark.asyncio
async def test_pkce_token_exchange_http200_no_error_field_no_access_token(self):
"""HTTP 200 with no error field and no access_token raises ProxyException
with a descriptive message showing the actual response keys."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
body_without_token = {"token_type": "Bearer", "scope": "openid"}
with patch("litellm.proxy.management_endpoints.ui_sso.httpx.AsyncClient") as mock_client_cls:
mock_client = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.json.return_value = body_without_token
mock_client.post = AsyncMock(return_value=mock_resp)
mock_client_cls.return_value = mock_client
with pytest.raises(ProxyException) as exc_info:
await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code="some_code",
code_verifier="verifier",
client_id="cid",
client_secret="csecret",
token_endpoint="https://example.com/token",
userinfo_endpoint=None,
include_client_id=False,
redirect_url=None,
additional_headers={},
)
assert "no access_token" in exc_info.value.message or "access_token" in exc_info.value.message
assert str(exc_info.value.code) == "401"
# Tests for SSO user team assignment bug (Issue: SSO Users Not Added to Entra-Synced Teams on First Login)
@ -4491,4 +5160,5 @@ def test_generic_response_convertor_extra_attributes_missing_field(monkeypatch):
assert result.extra_fields is not None
assert result.extra_fields["missing_field"] is None
assert result.extra_fields["another_missing"] is None
assert result.extra_fields["another_missing"] is None