fix(sso): replace httpx.AsyncClient() with get_async_httpx_client

Use the cached SSO_HANDLER client instead of creating a new
httpx.AsyncClient per request in PKCE token exchange and userinfo
fetch. Converts httpx.BasicAuth to a manual Authorization header
since AsyncHTTPHandler.post() does not accept an auth param.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
yuneng-jiang 2026-03-12 23:48:34 -07:00
parent 06681ddfcc
commit 2a997993d4

View file

@ -17,7 +17,6 @@ import secrets
from copy import deepcopy from copy import deepcopy
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast
import httpx
import jwt import jwt
from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi.responses import RedirectResponse from fastapi.responses import RedirectResponse
@ -2801,20 +2800,19 @@ class SSOAuthenticationHandler:
if redirect_url: if redirect_url:
token_data["redirect_uri"] = redirect_url token_data["redirect_uri"] = redirect_url
post_kwargs: Dict[str, Any] = { request_headers = {
"data": token_data,
"headers": {
**additional_headers, **additional_headers,
"Content-Type": "application/x-www-form-urlencoded", # must not be overridden "Content-Type": "application/x-www-form-urlencoded", # must not be overridden
"Accept": "application/json", "Accept": "application/json",
},
"timeout": 30.0,
} }
if not include_client_id: if not include_client_id:
# Use Basic Auth only when a secret is available; public PKCE clients omit it. # Use Basic Auth only when a secret is available; public PKCE clients omit it.
if client_secret: if client_secret:
post_kwargs["auth"] = httpx.BasicAuth(client_id, client_secret) credentials = base64.b64encode(
f"{client_id}:{client_secret}".encode()
).decode()
request_headers["Authorization"] = f"Basic {credentials}"
else: else:
token_data["client_id"] = client_id token_data["client_id"] = client_id
else: else:
@ -2822,13 +2820,16 @@ class SSOAuthenticationHandler:
if client_secret: if client_secret:
token_data["client_secret"] = client_secret token_data["client_secret"] = client_secret
# The try/except is INSIDE the async with so that TLS teardown exceptions http_client = get_async_httpx_client(
# from __aexit__ propagate as-is and are NOT mis-labelled as "Token endpoint llm_provider=httpxSpecialProvider.SSO_HANDLER
# 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: try:
response = await http_client.post(token_endpoint, **post_kwargs) response = await http_client.post(
url=token_endpoint,
data=token_data,
headers=request_headers,
timeout=30.0,
)
except Exception as exc: except Exception as exc:
# Catch network-level errors (SSL, DNS, TCP, timeout, etc.) and # Catch network-level errors (SSL, DNS, TCP, timeout, etc.) and
# wrap them as a clean ProxyException rather than leaking raw # wrap them as a clean ProxyException rather than leaking raw
@ -2840,9 +2841,6 @@ class SSOAuthenticationHandler:
param="token_exchange", param="token_exchange",
code=status.HTTP_401_UNAUTHORIZED, code=status.HTTP_401_UNAUTHORIZED,
) from exc ) 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: if response.status_code != 200:
verbose_proxy_logger.error( verbose_proxy_logger.error(
"PKCE token exchange failed. status=%s body=%s", "PKCE token exchange failed. status=%s body=%s",
@ -2970,9 +2968,11 @@ class SSOAuthenticationHandler:
if userinfo_endpoint: if userinfo_endpoint:
try: try:
async with httpx.AsyncClient() as client: client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.SSO_HANDLER
)
resp = await client.get( resp = await client.get(
userinfo_endpoint, url=userinfo_endpoint,
headers={ headers={
**additional_headers, **additional_headers,
"Authorization": f"Bearer {access_token}", # must not be overridden "Authorization": f"Bearer {access_token}", # must not be overridden