From 544ea04dc42dba6419fdbc3e930c3ff143346c7e Mon Sep 17 00:00:00 2001 From: Chris Withers Date: Tue, 14 Oct 2025 08:43:47 +0100 Subject: [PATCH] WIP - aiohttp: support for auth from ClientSession --- src/pook/interceptors/aiohttp.py | 56 +++++++++++++++++++++---- tests/unit/interceptors/aiohttp_test.py | 16 +++++++ 2 files changed, 65 insertions(+), 7 deletions(-) diff --git a/src/pook/interceptors/aiohttp.py b/src/pook/interceptors/aiohttp.py index 5646513..df7688c 100644 --- a/src/pook/interceptors/aiohttp.py +++ b/src/pook/interceptors/aiohttp.py @@ -9,6 +9,7 @@ from aiohttp.helpers import TimerNoop from aiohttp.streams import EmptyStreamReader +from pook.headers import HTTPHeaderDict from pook.request import Request # type: ignore from pook.interceptors.base import BaseInterceptor @@ -58,7 +59,8 @@ def set_headers(self, req, headers) -> None: # ``pook.request`` only allows a dict, so we need to map the iterable to the matchable interface if headers: if isinstance(headers, Mapping): - req.headers.update(**headers) + for key, val in headers.items(): + req.headers.add(key, val) else: # If it isn't a mapping, then its an Iterable[Tuple[Union[str, istr], str]] for req_header, req_header_value in headers: @@ -81,22 +83,20 @@ async def _on_request( # Create request contract based on incoming params req = Request(method) - self.set_headers(req, headers) - self.set_headers(req, session.headers) - - req.body = data + req.body = data # XXX take from real request? # Expose extra variadic arguments req.extra = kw full_url = session._build_url(url) + original_params = kw.get("params") # Compose URL - if not kw.get("params"): + if not original_params: req.url = str(full_url) else: # Transform params as a list of tuple - params = kw["params"] + params = original_params if isinstance(params, dict): params = [(x, y) for x, y in kw["params"].items()] req.url = str(full_url) + "?" + urlencode(params) @@ -107,6 +107,48 @@ async def _on_request( if "Content-Type" not in req.headers: req.headers["Content-Type"] = "application/json" + # Lifted from the ClientSession._request method we're mocking: + auth = kw.get('auth') + if ( + auth is None + and session._default_auth + and ( + not session._base_url or session._base_url_origin == full_url.origin() + ) + ): + auth = session._default_auth + + # Lifted from the ClientSession._request method we're mocking: + headers = session._prepare_headers(headers) + + aiohttp_req = session._request_class( + method, + full_url, + params=original_params, + headers=headers, + skip_auto_headers=None, # XXX + data=data, + cookies=None, # XXX + auth=auth, + version=session._version, + compress=kw.get('compress'), + chunked=kw.get('chunked'), + expect100=kw.get('expect100', False), + loop=session._loop, + response_class=session._response_class, + proxy=None, # XXX + proxy_auth=kw.get('proxy_auth'), + timer=None, # XXX, + session=session, + ssl=None, # XXX, + server_hostname=kw.get('server_hostname'), + proxy_headers=None, # XXX + traces=None, # XXX + trust_env=session.trust_env, + ) + + self.set_headers(req, aiohttp_req.headers) + # Match the request against the registered mocks in pook mock = self.engine.match(req) diff --git a/tests/unit/interceptors/aiohttp_test.py b/tests/unit/interceptors/aiohttp_test.py index c4bb193..f63a1ab 100644 --- a/tests/unit/interceptors/aiohttp_test.py +++ b/tests/unit/interceptors/aiohttp_test.py @@ -1,5 +1,6 @@ import aiohttp import pytest +from aiohttp import BasicAuth import pook from tests.unit.fixtures import BINARY_FILE @@ -97,6 +98,21 @@ async def test_client_headers_merged(local_responder): assert await res.read() == b"hello from pook" +@pytest.mark.asyncio +async def test_client_auth_merged(local_responder): + """Auth headers set on the client should be matched""" + pook \ + .get(local_responder + "/status/404") \ + .header("Authorization", "Basic dXNlcjpwYXNzd29yZA==") \ + .reply(200).body("hello from pook") + async with aiohttp.ClientSession(auth=BasicAuth('user', 'password')) as session: + res = await session.get( + local_responder + "/status/404", headers={"x-pook-secondary": "xyz"} + ) + assert res.status == 200 + assert await res.read() == b"hello from pook" + + @pytest.mark.asyncio async def test_client_headers_both_session_and_request(local_responder): """Headers should be matchable from both the session and request in the same matcher"""