Skip to content

Commit 7855de8

Browse files
authored
Merge pull request #16 from modern-python/story/2-4-auth-coercion
feat(story-2.4): auth coercion as middleware
2 parents b6eb1f7 + 17a8b4e commit 7855de8

10 files changed

Lines changed: 1789 additions & 6 deletions

docs/superpowers/plans/2026-06-01-auth-coercion-plan.md

Lines changed: 1006 additions & 0 deletions
Large diffs are not rendered by default.

docs/superpowers/specs/2026-06-01-auth-coercion-design.md

Lines changed: 372 additions & 0 deletions
Large diffs are not rendered by default.

src/httpware/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
"""httpware — resilience-first async HTTP client framework for Python."""
22

3+
from httpware._internal.auth import AuthValue
34
from httpware.client import AsyncClient
45
from httpware.config import ClientConfig, Limits, Timeout
56
from httpware.decoders import ResponseDecoder
@@ -33,6 +34,7 @@
3334
__all__ = [
3435
"STATUS_TO_EXCEPTION",
3536
"AsyncClient",
37+
"AuthValue",
3638
"BadRequestError",
3739
"ClientConfig",
3840
"ClientError",

src/httpware/_internal/auth.py

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
1+
"""Normalize the `auth=` value of AsyncClient into a Middleware (or None)."""
2+
3+
import inspect
4+
from collections.abc import Awaitable, Callable
5+
from typing import TypeAlias
6+
7+
from httpware.middleware import Middleware, before_request
8+
from httpware.request import Request
9+
10+
11+
_MIDDLEWARE_ARITY = 2
12+
13+
AuthValue: TypeAlias = str | Callable[[], str | Awaitable[str]] | Middleware | None
14+
15+
16+
def _normalize_auth(value: AuthValue) -> Middleware | None:
17+
"""Coerce an `auth=` value into a Middleware.
18+
19+
- `None` → returns `None` (no auth middleware injected).
20+
- `str` → returns a middleware that sets `Authorization: Bearer <str>`
21+
on every request (skipping if Authorization is already present).
22+
- `Callable[[], str | Awaitable[str]]` (zero-arg) → returns a middleware
23+
that calls the provider per request (awaiting if it returns an
24+
awaitable) and sets `Authorization: Bearer <result>` (skip-if-present).
25+
- `Middleware` (two-arg `__call__(request, next)`) → returned unchanged.
26+
- Any other callable shape → raises `TypeError` naming `auth=`.
27+
"""
28+
if value is None:
29+
return None
30+
if isinstance(value, str):
31+
return _bearer(value)
32+
if not callable(value):
33+
msg = f"`auth=` must be a string, zero-arg callable, Middleware, or None; got {type(value).__name__}"
34+
raise TypeError(msg)
35+
n_params = len(inspect.signature(value).parameters)
36+
if n_params == 0:
37+
return _bearer_from_provider(value) # ty: ignore[invalid-argument-type]
38+
if n_params == _MIDDLEWARE_ARITY:
39+
return value # ty: ignore[invalid-return-type]
40+
msg = f"`auth=` callable must take 0 args (token provider) or 2 args (Middleware); got {n_params}"
41+
raise TypeError(msg)
42+
43+
44+
def _bearer(token: str) -> Middleware:
45+
"""Middleware that sets `Authorization: Bearer <token>` (skip-if-present)."""
46+
47+
@before_request
48+
async def _add_static_bearer(request: Request) -> Request:
49+
if _has_authorization(request):
50+
return request
51+
return request.with_header("Authorization", f"Bearer {token}")
52+
53+
return _add_static_bearer
54+
55+
56+
def _bearer_from_provider(
57+
provider: Callable[[], str | Awaitable[str]],
58+
) -> Middleware:
59+
"""Middleware that calls `provider()` per request and sets the header."""
60+
61+
@before_request
62+
async def _add_dynamic_bearer(request: Request) -> Request:
63+
if _has_authorization(request):
64+
return request
65+
token = provider()
66+
if inspect.isawaitable(token):
67+
token = await token
68+
return request.with_header("Authorization", f"Bearer {token}")
69+
70+
return _add_dynamic_bearer
71+
72+
73+
def _has_authorization(request: Request) -> bool:
74+
"""Case-insensitive check for an existing Authorization header."""
75+
return any(k.lower() == "authorization" for k in request.headers)

src/httpware/client.py

Lines changed: 45 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
import typing
66
from collections.abc import Mapping, Sequence
77

8+
from httpware._internal.auth import AuthValue, _normalize_auth
89
from httpware._internal.chain import compose
910
from httpware.config import ClientConfig, Limits, Timeout
1011
from httpware.decoders import ResponseDecoder
@@ -52,6 +53,8 @@ class AsyncClient:
5253
_transport: Transport
5354
_dispatch: Next
5455
_owns_transport: bool
56+
_user_middleware: tuple[Middleware, ...]
57+
_auth: AuthValue
5558

5659
def __init__(
5760
self,
@@ -64,12 +67,19 @@ def __init__(
6467
transport: Transport | None = None,
6568
decoder: ResponseDecoder | None = None,
6669
middleware: Sequence[Middleware] | None = None,
70+
auth: AuthValue = None,
6771
) -> None:
6872
normalized_timeout = _normalize_timeout(timeout)
6973
resolved_limits = limits or Limits()
7074
resolved_transport: Transport = transport or Httpx2Transport(limits=resolved_limits, timeout=normalized_timeout)
7175
resolved_decoder = decoder or PydanticDecoder()
72-
resolved_middleware = tuple(middleware) if middleware is not None else ()
76+
resolved_user_middleware: tuple[Middleware, ...] = tuple(middleware) if middleware is not None else ()
77+
resolved_auth_middleware = _normalize_auth(auth)
78+
composed_middleware: tuple[Middleware, ...] = (
79+
resolved_user_middleware
80+
if resolved_auth_middleware is None
81+
else (*resolved_user_middleware, resolved_auth_middleware)
82+
)
7383

7484
self._config = ClientConfig(
7585
base_url=base_url,
@@ -78,11 +88,13 @@ def __init__(
7888
timeout=normalized_timeout,
7989
limits=resolved_limits,
8090
decoder=resolved_decoder,
81-
middleware=resolved_middleware,
91+
middleware=composed_middleware,
8292
)
8393
self._transport = resolved_transport
84-
self._dispatch = compose(resolved_middleware, resolved_transport)
94+
self._dispatch = compose(composed_middleware, resolved_transport)
8595
self._owns_transport = True
96+
self._user_middleware = resolved_user_middleware
97+
self._auth = auth
8698

8799
@classmethod
88100
def from_url(cls, base_url: str, **kwargs: object) -> "AsyncClient":
@@ -582,6 +594,7 @@ def with_options(
582594
timeout: Timeout | float | None = _UNSET,
583595
decoder: ResponseDecoder | None = _UNSET,
584596
middleware: Sequence[Middleware] | None = _UNSET,
597+
auth: AuthValue | object = _UNSET,
585598
) -> "AsyncClient":
586599
"""Return a new AsyncClient sharing the same transport with overridden config.
587600
@@ -603,18 +616,44 @@ def with_options(
603616
changes["timeout"] = _normalize_timeout(timeout)
604617
if decoder is not _UNSET:
605618
changes["decoder"] = decoder or PydanticDecoder()
619+
620+
new_user_middleware = self._user_middleware
606621
if middleware is not _UNSET:
607-
changes["middleware"] = tuple(middleware) if middleware is not None else ()
622+
new_user_middleware = tuple(middleware) if middleware is not None else ()
623+
624+
new_auth: AuthValue = self._auth
625+
if auth is not _UNSET:
626+
new_auth = auth # ty: ignore[invalid-assignment]
627+
628+
new_auth_middleware = _normalize_auth(new_auth)
629+
new_composed: tuple[Middleware, ...] = (
630+
new_user_middleware if new_auth_middleware is None else (*new_user_middleware, new_auth_middleware)
631+
)
632+
changes["middleware"] = new_composed
608633

609634
new_config = dataclasses.replace(self._config, **changes)
610-
return AsyncClient._from_view(new_config, self._transport)
635+
return AsyncClient._from_view(
636+
new_config,
637+
self._transport,
638+
user_middleware=new_user_middleware,
639+
auth=new_auth,
640+
)
611641

612642
@classmethod
613-
def _from_view(cls, config: ClientConfig, transport: Transport) -> "AsyncClient":
643+
def _from_view(
644+
cls,
645+
config: ClientConfig,
646+
transport: Transport,
647+
*,
648+
user_middleware: tuple[Middleware, ...],
649+
auth: AuthValue,
650+
) -> "AsyncClient":
614651
"""Construct a view sharing an existing transport. Bypasses __init__."""
615652
client = cls.__new__(cls)
616653
client._config = config # noqa: SLF001
617654
client._transport = transport # noqa: SLF001
618655
client._dispatch = compose(config.middleware, transport) # noqa: SLF001
619656
client._owns_transport = False # noqa: SLF001
657+
client._user_middleware = user_middleware # noqa: SLF001
658+
client._auth = auth # noqa: SLF001
620659
return client

tests/test_client_construction.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -76,3 +76,37 @@ def test_construction_does_not_create_httpx2_client() -> None:
7676
# Httpx2Transport stores `_client` lazily; until first call, _client is None.
7777
# The attribute is private; we check it via getattr to keep the test resilient.
7878
assert getattr(client._transport, "_client", "missing") is None
79+
80+
81+
def test_init_no_auth_means_no_auth_middleware() -> None:
82+
transport = RecordedTransport()
83+
client = AsyncClient(transport=transport)
84+
assert client._config.middleware == ()
85+
assert client._auth is None
86+
assert client._user_middleware == ()
87+
88+
89+
def test_init_with_string_auth_appends_bearer_middleware() -> None:
90+
transport = RecordedTransport()
91+
client = AsyncClient(transport=transport, auth="tok")
92+
assert len(client._config.middleware) == 1
93+
assert isinstance(client._config.middleware[0], Middleware)
94+
assert client._auth == "tok"
95+
assert client._user_middleware == ()
96+
97+
98+
def test_init_with_user_middleware_plus_auth() -> None:
99+
class _M:
100+
async def __call__(self, request, next) -> Response: # noqa: A002, ANN001
101+
return await next(request)
102+
103+
m1 = _M()
104+
m2 = _M()
105+
transport = RecordedTransport()
106+
client = AsyncClient(transport=transport, middleware=[m1, m2], auth="tok")
107+
_expected_len = 3
108+
assert len(client._config.middleware) == _expected_len
109+
assert client._config.middleware[0] is m1
110+
assert client._config.middleware[1] is m2
111+
# The third entry is the auth middleware; identity-test that user_middleware excludes it.
112+
assert client._user_middleware == (m1, m2)

tests/test_client_methods.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -152,3 +152,40 @@ async def test_per_call_timeout_propagates_to_request_extensions() -> None:
152152
await client.get("/foo", timeout=2.5)
153153
assert transport.last_request is not None
154154
assert "timeout" in transport.last_request.extensions
155+
156+
157+
async def test_string_auth_sends_authorization_header() -> None:
158+
transport = RecordedTransport(default=Response(status=200, headers={}, content=b"", url="/", elapsed=0.0))
159+
client = AsyncClient(transport=transport, auth="tok")
160+
161+
await client.get("/foo")
162+
163+
assert transport.last_request is not None
164+
assert transport.last_request.headers["Authorization"] == "Bearer tok"
165+
166+
167+
async def test_per_call_authorization_header_wins_over_auth_param() -> None:
168+
transport = RecordedTransport(default=Response(status=200, headers={}, content=b"", url="/", elapsed=0.0))
169+
client = AsyncClient(transport=transport, auth="default-tok")
170+
171+
await client.get("/foo", headers={"Authorization": "Bearer override"})
172+
173+
assert transport.last_request is not None
174+
assert transport.last_request.headers["Authorization"] == "Bearer override"
175+
176+
177+
async def test_callable_auth_calls_provider_per_request() -> None:
178+
transport = RecordedTransport(default=Response(status=200, headers={}, content=b"", url="/", elapsed=0.0))
179+
calls = 0
180+
181+
def _provider() -> str:
182+
nonlocal calls
183+
calls += 1
184+
return f"tok-{calls}"
185+
186+
client = AsyncClient(transport=transport, auth=_provider)
187+
188+
await client.get("/a")
189+
await client.get("/b")
190+
191+
assert calls == 2 # noqa: PLR2004

tests/test_client_middleware_wiring.py

Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
"""Unit tests for AsyncClient middleware wiring through compose() and with_options."""
22

3+
from collections.abc import Mapping
4+
35
from httpware import AsyncClient, RecordedTransport
46
from httpware.middleware import Middleware, Next
57
from httpware.request import Request
@@ -112,3 +114,54 @@ def decode(self, content: bytes, model: type) -> object: # pragma: no cover #
112114
client = AsyncClient(transport=transport)
113115
view = client.with_options(decoder=new_decoder)
114116
assert view._config.decoder is new_decoder # noqa: SLF001
117+
118+
119+
async def test_auth_runs_inside_user_middleware() -> None:
120+
transport = RecordedTransport(default=Response(status=200, headers={}, content=b"", url="/", elapsed=0.0))
121+
122+
user_seen_headers: list[Mapping[str, str]] = []
123+
124+
class _UserOuter:
125+
async def __call__(self, request: Request, next: Next) -> Response: # noqa: A002
126+
user_seen_headers.append(dict(request.headers))
127+
return await next(request)
128+
129+
client = AsyncClient(transport=transport, middleware=[_UserOuter()], auth="tok")
130+
await client.get("/foo")
131+
132+
# User middleware saw the request BEFORE auth header was applied.
133+
assert "Authorization" not in user_seen_headers[0]
134+
# Transport saw the request WITH the auth header.
135+
assert transport.last_request is not None
136+
assert transport.last_request.headers["Authorization"] == "Bearer tok"
137+
138+
139+
async def test_with_options_auth_replaces_auth_middleware() -> None:
140+
transport = RecordedTransport(default=Response(status=200, headers={}, content=b"", url="/", elapsed=0.0))
141+
client = AsyncClient(transport=transport, auth="parent")
142+
view = client.with_options(auth="view")
143+
144+
await view.get("/foo")
145+
assert transport.last_request is not None
146+
assert transport.last_request.headers["Authorization"] == "Bearer view"
147+
148+
await client.get("/foo")
149+
assert transport.last_request is not None
150+
assert transport.last_request.headers["Authorization"] == "Bearer parent"
151+
152+
153+
async def test_with_options_middleware_keeps_existing_auth() -> None:
154+
transport = RecordedTransport(default=Response(status=200, headers={}, content=b"", url="/", elapsed=0.0))
155+
156+
class _M:
157+
async def __call__(self, request: Request, next: Next) -> Response: # noqa: A002
158+
return await next(request)
159+
160+
m1 = _M()
161+
m2 = _M()
162+
client = AsyncClient(transport=transport, auth="tok", middleware=[m1])
163+
view = client.with_options(middleware=[m2])
164+
165+
await view.get("/foo")
166+
assert transport.last_request is not None
167+
assert transport.last_request.headers["Authorization"] == "Bearer tok"

0 commit comments

Comments
 (0)