From f8c8419d9eec9ccb95ccd87174459879152973fc Mon Sep 17 00:00:00 2001 From: Jacob Henner Date: Fri, 25 Sep 2026 17:00:18 -0400 Subject: [PATCH] Differentiate async auth Differentiate async auth, to ensure that any blocking IO performed during GSSAPI authentication (e.g. network interactions with a KDC) does not block the event loop. --- httpx_gssapi/gssapi_.py | 82 +++++++++++++++++++++++++++++------------ tests/test_mocked.py | 59 +++++++++++++++++++++++++++-- 2 files changed, 115 insertions(+), 26 deletions(-) diff --git a/httpx_gssapi/gssapi_.py b/httpx_gssapi/gssapi_.py index 0f5a99b..b477ac6 100644 --- a/httpx_gssapi/gssapi_.py +++ b/httpx_gssapi/gssapi_.py @@ -1,8 +1,9 @@ +import asyncio import re import logging from itertools import chain from functools import wraps -from typing import Generator, Optional, List, Any, Union +from typing import AsyncGenerator, Generator, Optional, List, Any, Union from base64 import b64encode, b64decode @@ -17,6 +18,7 @@ log = logging.getLogger(__name__) FlowGen = Generator[Request, Response, None] +AsyncFlowGen = AsyncGenerator[Request, Response] # Different types of mutual authentication: # with mutual_authentication set to REQUIRED, all responses will be @@ -98,8 +100,7 @@ def _gss_to_spnego_error(message: str, *args: Any, **kwargs: Any): """Helper function to _handle_gsserror to raise SPNEGOExchangeErrors.""" try: request = next( - a for a in chain(args, kwargs.values()) - if isinstance(a, Request) + a for a in chain(args, kwargs.values()) if isinstance(a, Request) ) except StopIteration: # sanity check raise RuntimeError("No request in arguments!") @@ -136,14 +137,16 @@ class HTTPSPNEGOAuth(Auth): """ - def __init__(self, - mutual_authentication: int = DISABLED, - target_name: Optional[Union[str, gssapi.Name]] = "HTTP", - delegate: bool = False, - opportunistic_auth: bool = False, - creds: Optional[gssapi.Credentials] = None, - mech: Optional[Union[bytes, gssapi.OID]] = SPNEGO, - sanitize_mutual_error_response: bool = True): + def __init__( + self, + mutual_authentication: int = DISABLED, + target_name: Optional[Union[str, gssapi.Name]] = "HTTP", + delegate: bool = False, + opportunistic_auth: bool = False, + creds: Optional[gssapi.Credentials] = None, + mech: Optional[Union[bytes, gssapi.OID]] = SPNEGO, + sanitize_mutual_error_response: bool = True, + ): self.mutual_authentication = mutual_authentication self.target_name = target_name self.delegate = delegate @@ -152,7 +155,7 @@ def __init__(self, self.mech = mech self.sanitize_mutual_error_response = sanitize_mutual_error_response - def auth_flow(self, request: Request) -> FlowGen: + def sync_auth_flow(self, request: Request) -> FlowGen: if self.opportunistic_auth: # add Authorization header before we receive a 401 ctx = self.set_auth_header(request) @@ -162,9 +165,42 @@ def auth_flow(self, request: Request) -> FlowGen: response = yield request yield from self.handle_response(response, ctx) - def handle_response(self, - response: Response, - ctx: Optional[SecurityContext] = None) -> FlowGen: + async def async_auth_flow(self, request: Request) -> AsyncFlowGen: + if self.opportunistic_auth: + # GSSAPI may perform blocking I/O while acquiring credentials. + ctx = await asyncio.to_thread(self.set_auth_header, request) + else: + ctx = None + + response = yield request + num_401s = 0 + while response.status_code == 401 and num_401s < 2: + num_401s += 1 + log.debug("Handling 401 response, total seen: %d", num_401s) + + if _negotiate_value(response) is None: + log.debug("GSSAPI is not supported") + break + + log.debug("Generating user authentication header") + try: + ctx = await asyncio.to_thread( + self.set_auth_header, response.request, response + ) + except SPNEGOExchangeError: + log.debug("Failed to generate authentication header") + + response = yield response.request + + if response.status_code == 401 or ctx is None: + log.debug("Failed to authenticate, returning 401 response") + return + + await asyncio.to_thread(self.handle_mutual_auth, response, ctx) + + def handle_response( + self, response: Response, ctx: Optional[SecurityContext] = None + ) -> FlowGen: num_401s = 0 while response.status_code == 401 and num_401s < 2: num_401s += 1 @@ -221,8 +257,10 @@ def handle_mutual_auth(self, response: Response, ctx: SecurityContext): " on %d response", response.status_code, ) - if (self.mutual_authentication == REQUIRED - and self.sanitize_mutual_error_response): + if ( + self.mutual_authentication == REQUIRED + and self.sanitize_mutual_error_response + ): _sanitize_response(response) else: # Unable to attempt mutual authentication when mutual auth is @@ -232,9 +270,9 @@ def handle_mutual_auth(self, response: Response, ctx: SecurityContext): raise MutualAuthenticationError(response=response) @_handle_gsserror(gss_stage='stepping', result=_gss_to_spnego_error) - def set_auth_header(self, - request: Request, - response: Optional[Response] = None) -> SecurityContext: + def set_auth_header( + self, request: Request, response: Optional[Response] = None + ) -> SecurityContext: """ Create a new security context, generate the GSSAPI authentication token, and insert it into the request header. The new security context @@ -257,9 +295,7 @@ def set_auth_header(self, return ctx @_handle_gsserror(gss_stage="stepping", result=False) - def authenticate_server(self, - response: Response, - ctx: SecurityContext) -> bool: + def authenticate_server(self, response: Response, ctx: SecurityContext) -> bool: """ Uses GSSAPI to authenticate the server by extracting the negotiate value from the response and stepping the security context. diff --git a/tests/test_mocked.py b/tests/test_mocked.py index 770bfc3..5500331 100644 --- a/tests/test_mocked.py +++ b/tests/test_mocked.py @@ -1,7 +1,9 @@ #!/usr/bin/env python """Tests for httpx_gssapi.""" +import asyncio import logging +import threading from base64 import b64encode from unittest.mock import Mock, patch @@ -103,7 +105,7 @@ def test_force_preemptive(patched_ctx): request = null_request() - flow = auth.auth_flow(request) + flow = auth.sync_auth_flow(request) next(flow) # Move to first request yield assert 'Authorization' in request.headers @@ -115,7 +117,7 @@ def test_no_force_preemptive(patched_ctx): request = null_request() - flow = auth.auth_flow(request) + flow = auth.sync_auth_flow(request) next(flow) # Move to first request yield assert 'Authorization' not in request.headers @@ -363,13 +365,64 @@ def test_opportunistic_auth(patched_ctx): request = null_request() - flow = auth.auth_flow(request) + flow = auth.sync_auth_flow(request) assert next(flow) is request assert 'Authorization' in request.headers assert request.headers.get('Authorization') == b64_negotiate_response +def test_async_auth_flow_offloads_gssapi_work(): + auth = httpx_gssapi.HTTPSPNEGOAuth(opportunistic_auth=True) + event_loop_thread = threading.get_ident() + gssapi_thread = None + + def set_auth_header(request, response=None): + nonlocal gssapi_thread + gssapi_thread = threading.get_ident() + request.headers['Authorization'] = b64_negotiate_response + return object() + + async def request(): + transport = httpx.MockTransport( + lambda request: httpx.Response(200, request=request) + ) + with patch.object(auth, 'set_auth_header', side_effect=set_auth_header): + async with httpx.AsyncClient(auth=auth, transport=transport) as client: + response = await client.get('http://www.example.org/') + assert response.status_code == 200 + + asyncio.run(request()) + assert gssapi_thread is not None + assert gssapi_thread != event_loop_thread + + +def test_async_auth_flow_handles_negotiate_challenge(patched_ctx): + auth = httpx_gssapi.HTTPSPNEGOAuth() + request_count = 0 + + def transport(request): + nonlocal request_count + request_count += 1 + if request_count == 1: + return httpx.Response(401, headers=neg_token, request=request) + assert request.headers['Authorization'] == b64_negotiate_response + return httpx.Response(200, request=request) + + async def request(): + async with httpx.AsyncClient( + auth=auth, + transport=httpx.MockTransport(transport), + ) as client: + response = await client.get('http://www.example.org/') + assert response.status_code == 200 + + asyncio.run(request()) + assert request_count == 2 + check_init() + fake_resp.assert_called_with(b'token') + + def test_explicit_creds(patched_creds, patched_ctx): response = null_response(headers=neg_token) creds = gssapi.Credentials()