diff --git a/src/atproto/oauth.py b/src/atproto/oauth.py index 237e311..a8acc3b 100644 --- a/src/atproto/oauth.py +++ b/src/atproto/oauth.py @@ -26,6 +26,7 @@ class OAuthTokens(NamedTuple): # Prepares and sends a pushed auth request (PAR) via HTTP POST to the Authorization Server. # Returns "state" id HTTP response on success, without checking HTTP response status async def send_par_auth_request( + hardened_client: ClientSession, authserver_url: str, authserver_meta: dict[str, str], login_hint: str | None, @@ -70,16 +71,15 @@ async def send_par_auth_request( # IMPORTANT: Pushed Authorization Request URL is untrusted input, SSRF mitigations are needed assert is_safe_url(par_url) - async with hardened_http.get_session() as session: - resp = await session.post( - par_url, - headers={ - "Content-Type": "application/x-www-form-urlencoded", - "DPoP": dpop_proof, - }, - data=par_body, - ) - respjson = await resp.json() + resp = await hardened_client.post( + par_url, + headers={ + "Content-Type": "application/x-www-form-urlencoded", + "DPoP": dpop_proof, + }, + data=par_body, + ) + respjson = await resp.json() # Handle DPoP missing/invalid nonce error by retrying with server-provided nonce if resp.status == 400 and respjson["error"] == "use_dpop_nonce": diff --git a/src/oauth.py b/src/oauth.py index ee4d3ff..cc81d06 100644 --- a/src/oauth.py +++ b/src/oauth.py @@ -29,7 +29,7 @@ from src.auth import ( save_auth_session, ) from src.db import KV, get_db -from src.security import is_safe_url +from src.security import hardened_http, is_safe_url oauth = Blueprint("oauth", __name__, url_prefix="/oauth") @@ -100,7 +100,9 @@ async def oauth_start(): CLIENT_SECRET_JWK = JsonWebKey.import_key(current_app.config["CLIENT_SECRET_JWK"]) + client = hardened_http.get_session() pkce_verifier, state, dpop_authserver_nonce, resp = await send_par_auth_request( + client, authserver_url, authserver_meta, login_hint, @@ -119,8 +121,9 @@ async def oauth_start(): respjson: dict[str, str] = await resp.json() par_request_uri: str = respjson["request_uri"] - current_app.logger.debug(f"saving oauth_auth_request to DB state={state}") + await client.close() + current_app.logger.debug(f"saving oauth_auth_request to DB state={state}") oauth_request = OAuthAuthRequest( state, authserver_meta["issuer"],