fix(sso): extract PKCE callback validation to fix ruff PLR0915

get_generic_sso_response had 52 statements (ruff PLR0915 limit is 50).
Move the 5 PKCE callback guards (state/code/client_id/token_endpoint/type)
into a new _validate_pkce_oauth_callback helper. Pure refactor — same
exceptions, same control flow.
This commit is contained in:
mateo-berri 2026-05-16 01:09:05 +00:00
parent 52090ba31d
commit 21a5388da0
No known key found for this signature in database

View file

@ -1167,6 +1167,73 @@ def _handle_generic_sso_error(
raise e
def _validate_pkce_oauth_callback(
request: Request,
authorization_code: Optional[str],
generic_client_id: str,
generic_token_endpoint: Optional[str],
provider: _OIDC_PROVIDER_NAMES,
) -> str:
"""Validate the OAuth callback for the PKCE branch and return the authorization code.
The non-PKCE branch delegates to fastapi-sso's ``verify_and_process``, which
performs its own session-cookie check. The PKCE branch bypasses that helper,
so we validate the URL ``state`` against the ``litellm_oauth_state`` cookie
set on the redirect response without this an attacker can pre-mint a state
+ cached PKCE verifier and hijack a victim's auth code (Login-CSRF / token
theft).
"""
url_state = request.query_params.get("state")
cookie_state = request.cookies.get("litellm_oauth_state")
if (
not url_state
or not cookie_state
or not secrets.compare_digest(url_state, cookie_state)
):
raise ProxyException(
message=(
"Invalid OAuth state parameter — does not match "
"the browser-bound state cookie."
),
type=ProxyErrorTypes.auth_error,
param="state",
code=status.HTTP_400_BAD_REQUEST,
)
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:
client_id_env = f"{_get_oidc_env_prefix(provider)}_CLIENT_ID"
raise ProxyException(
message=f"{client_id_env} must be set when PKCE is enabled",
type=ProxyErrorTypes.auth_error,
param=client_id_env,
code=status.HTTP_401_UNAUTHORIZED,
)
if not generic_token_endpoint:
token_endpoint_env = f"{_get_oidc_env_prefix(provider)}_TOKEN_ENDPOINT"
raise ProxyException(
message=(f"{token_endpoint_env} must be set when PKCE is enabled"),
type=ProxyErrorTypes.auth_error,
param=token_endpoint_env,
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,
)
return authorization_code
async def get_generic_sso_response(
request: Request,
jwt_handler: JWTHandler,
@ -1259,62 +1326,13 @@ async def get_generic_sso_response(
authorization_code = request.query_params.get("code")
if code_verifier:
# State-to-session-cookie binding. The non-PKCE branch below
# delegates to fastapi-sso's ``verify_and_process``, which
# performs its own session-cookie check. The PKCE branch
# bypasses that helper, so we validate the URL ``state``
# against the ``litellm_oauth_state`` cookie set on the
# redirect response — without this an attacker can pre-mint
# a state + cached PKCE verifier and hijack a victim's auth
# code (Login-CSRF / token theft).
url_state = request.query_params.get("state")
cookie_state = request.cookies.get("litellm_oauth_state")
if (
not url_state
or not cookie_state
or not secrets.compare_digest(url_state, cookie_state)
):
raise ProxyException(
message=(
"Invalid OAuth state parameter — does not match "
"the browser-bound state cookie."
),
type=ProxyErrorTypes.auth_error,
param="state",
code=status.HTTP_400_BAD_REQUEST,
)
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:
client_id_env = f"{_get_oidc_env_prefix(provider)}_CLIENT_ID"
raise ProxyException(
message=f"{client_id_env} must be set when PKCE is enabled",
type=ProxyErrorTypes.auth_error,
param=client_id_env,
code=status.HTTP_401_UNAUTHORIZED,
)
if not generic_token_endpoint:
token_endpoint_env = f"{_get_oidc_env_prefix(provider)}_TOKEN_ENDPOINT"
raise ProxyException(
message=(f"{token_endpoint_env} must be set when PKCE is enabled"),
type=ProxyErrorTypes.auth_error,
param=token_endpoint_env,
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,
)
authorization_code = _validate_pkce_oauth_callback(
request=request,
authorization_code=authorization_code,
generic_client_id=generic_client_id,
generic_token_endpoint=generic_token_endpoint,
provider=provider,
)
combined_response = await SSOAuthenticationHandler._pkce_token_exchange(
authorization_code=authorization_code,
code_verifier=code_verifier,