Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
82 changes: 59 additions & 23 deletions httpx_gssapi/gssapi_.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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
Expand Down Expand Up @@ -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!")
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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.
Expand Down
59 changes: 56 additions & 3 deletions tests/test_mocked.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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()
Expand Down