From 7697b1c3976c90c09920a8e0c0000dec2328faf0 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 12 Mar 2026 12:41:39 -0700 Subject: [PATCH] fix(sso): direct PKCE token exchange + Redis wiring for multi-instance SSO (#22923) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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 --- litellm/proxy/management_endpoints/ui_sso.py | 659 ++++++++++++++- litellm/proxy/proxy_server.py | 27 +- .../proxy/management_endpoints/test_ui_sso.py | 748 +++++++++++++++++- 3 files changed, 1355 insertions(+), 79 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 084a528e8a8..0161c488d05 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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: """ diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e01ee8e9d9e..77e2a88796e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 3a8d17ccb45..ca480b84420 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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 \ No newline at end of file + assert result.extra_fields["another_missing"] is None +