Skip to content

Commit 19e8b03

Browse files
fix(auth): sanitize HTTP headers and preserve empty responses
1 parent 8293ebe commit 19e8b03

8 files changed

Lines changed: 120 additions & 5 deletions

File tree

‎aws_lambda_powertools/utilities/auth_alpha/_internal/http.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from __future__ import annotations
22

33
import json
4+
from io import BytesIO
45
from typing import TYPE_CHECKING, Any, cast
56

67
import urllib3
@@ -109,7 +110,7 @@ def request(
109110
)
110111
content = self._read_body(response, deadline)
111112
return urllib3.HTTPResponse(
112-
body=content,
113+
body=BytesIO(content),
113114
status=response.status,
114115
headers=response.headers,
115116
reason=response.reason,

‎aws_lambda_powertools/utilities/auth_alpha/_internal/transport.py‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,10 @@
88

99
from urllib3.connection import HTTPSConnection
1010
from urllib3.connectionpool import HTTPSConnectionPool
11+
from urllib3.exceptions import HeaderParsingError
12+
from urllib3.util.response import assert_header_parsing
13+
14+
from aws_lambda_powertools.utilities.auth_alpha._internal.deadline import RequestError
1115

1216
if TYPE_CHECKING:
1317
from collections.abc import Iterator
@@ -77,6 +81,15 @@ def __init__(
7781
# The stream owns the socket reference even for Connection: close.
7882
self.fp = BufferedReader(_DeadlineReader(self.fp, sock, deadline), buffer_size=8192)
7983

84+
def begin(self) -> None:
85+
super().begin()
86+
try:
87+
# urllib3 otherwise logs malformed headers, including provider data,
88+
# before the public Auth operation can sanitize the failure.
89+
assert_header_parsing(self.msg)
90+
except (HeaderParsingError, TypeError):
91+
raise RequestError() from None
92+
8093

8194
class _DeadlineHTTPSConnection(HTTPSConnection):
8295
response_class = _DeadlineResponse

‎aws_lambda_powertools/utilities/auth_alpha/oauth2/client.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,8 @@
2828

2929
_BEARER_TOKEN = re.compile(r"[-A-Za-z0-9._~+/]+=*")
3030
_HEADER_NAME = re.compile(r"[-!#$%&'*+.^_`|~0-9A-Za-z]+")
31+
# HTTP field values allow horizontal tabs, visible ASCII, and extended Latin-1 bytes.
32+
_HEADER_VALUE = re.compile(r"[\t\x20-\x7e\x80-\xff]*")
3133
_RESOURCE_CHARACTERS = frozenset(ascii_letters + digits + "-._~:/?[]@!$&'()*+,;=")
3234
_HEX_DIGITS = frozenset(hexdigits)
3335

@@ -265,7 +267,7 @@ def _request_headers(headers: Mapping[str, str] | None) -> dict[str, str]:
265267
or not isinstance(value, str)
266268
or not _HEADER_NAME.fullmatch(name)
267269
or name.lower() == "authorization"
268-
or any(character in value for character in ("\r", "\n"))
270+
or not _HEADER_VALUE.fullmatch(value)
269271
):
270272
raise ValueError("Request headers must be valid and must not include Authorization")
271273
return dict(headers)

‎docs/utilities/oauth2.md‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -106,6 +106,8 @@ The helper requires HTTPS, rejects an existing Authorization header, and never f
106106

107107
Header names must use HTTP token syntax: letters, digits, and the permitted token punctuation. Empty names, whitespace (including trailing spaces or tabs), and delimiters such as colons are rejected before token acquisition. Authorization is rejected regardless of casing.
108108

109+
Header values must fit Latin-1 and cannot contain ASCII control characters other than horizontal tabs. Invalid names and values are rejected before loading the client secret.
110+
109111
!!! warning "Use trusted destination URLs"
110112
`request()` does not derive or restrict destinations from the configured audience or resource. Supply trusted URLs from application configuration; never pass a caller-controlled destination. A token intended for one API must not be sent to another.
111113

‎tests/functional/auth_alpha/oauth2/test_client.py‎

Lines changed: 22 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -704,11 +704,31 @@ def load_secret():
704704
assert http.requests == []
705705

706706

707-
def test_http_token_punctuation_is_allowed_in_header_names(http):
707+
@pytest.mark.parametrize("value", ["\x00", "\x01", "\x1f", "\x7f", "東京", "\ud800"])
708+
def test_invalid_header_values_are_rejected_before_token_acquisition(http, value):
709+
secrets = []
710+
711+
def load_secret():
712+
secrets.append("test-secret")
713+
return secrets[-1]
714+
715+
http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST")
716+
http.serve("https://api.example.com/orders", {"orders": []})
717+
subject = client(client_secret=load_secret)
718+
719+
with pytest.raises(ValueError, match="Request headers"):
720+
subject.request("GET", "https://api.example.com/orders", headers={"X-Trace": f"trace{value}value"})
721+
722+
assert secrets == []
723+
assert http.requests == []
724+
725+
726+
@pytest.mark.parametrize("value", ["", "trace-id", "trace\tvalue", "\x80", "\xff", "caf\xe9"])
727+
def test_valid_header_names_and_values_are_preserved(http, value):
708728
http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST")
709729
http.serve("https://api.example.com/orders", {"orders": []})
710730
subject = client()
711-
headers = {"X-Trace!#$%&'*+.^_`|~09": "trace-id"}
731+
headers = {"X-Trace!#$%&'*+.^_`|~09": value}
712732

713733
response = subject.request("GET", "https://api.example.com/orders", headers=headers)
714734

‎tests/integration/auth_alpha/conftest.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -156,6 +156,9 @@ def do_GET(self): # noqa: N802
156156
def do_POST(self): # noqa: N802
157157
self.respond()
158158

159+
def do_HEAD(self): # noqa: N802
160+
self.respond()
161+
159162
def respond(self):
160163
body = self.rfile.read(int(self.headers.get("Content-Length", 0)))
161164
endpoint.requests.append((self.command, self.path, dict(self.headers), body))
@@ -180,7 +183,8 @@ def respond(self):
180183
):
181184
return
182185
self.end_headers()
183-
_write_body(self.wfile, reply, endpoint.stop)
186+
if self.command != "HEAD":
187+
_write_body(self.wfile, reply, endpoint.stop)
184188
except (OSError, ssl.SSLError):
185189
# Timeout and oversized-body tests deliberately close early.
186190
pass

‎tests/integration/auth_alpha/jwt/test_https.py‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import time
2+
import traceback
23

34
import jwt
45
import pytest
@@ -70,3 +71,19 @@ def test_key_endpoint_failures_are_bounded_and_do_not_follow_redirects(https_ser
7071
assert time.monotonic() - started < 1
7172
assert error.value.__context__ is None
7273
assert [request[1] for request in https_server.requests] == ["/keys"]
74+
75+
76+
def test_malformed_jwks_headers_fail_without_logging_provider_data(https_server, caplog):
77+
private_data = "local-test-private-provider-data"
78+
https_server.serve("/keys", {"keys": []}, headers={"Broken header": private_data})
79+
subject = verifier(https_server, jwks_uri=https_server.url + "/keys")
80+
81+
with pytest.raises(JWKSFetchError) as error:
82+
subject.prefetch()
83+
84+
assert error.value.__context__ is None
85+
assert error.value.__cause__ is None
86+
assert private_data not in "".join(traceback.format_exception(error.value))
87+
assert private_data not in caplog.text
88+
assert not [record for record in caplog.records if record.name == "urllib3.connection"]
89+
assert [request[1] for request in https_server.requests] == ["/keys"]

‎tests/integration/auth_alpha/oauth2/test_https.py‎

Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import base64
22
import threading
33
import time
4+
import traceback
45
from concurrent.futures import ThreadPoolExecutor
56
from urllib.parse import parse_qs
67

@@ -223,3 +224,58 @@ def test_downstream_gzip_response_remains_readable(https_server, chunked):
223224

224225
response = subject.request("GET", https_server.url + "/inventory")
225226
assert response.json() == {"items": [123]}
227+
228+
229+
@pytest.mark.parametrize(("method", "status"), [("GET", 200), ("GET", 204), ("GET", 304), ("HEAD", 200)])
230+
def test_empty_downstream_responses_preserve_bytes_and_status(https_server, method, status):
231+
https_server.serve("/token", TOKEN_RESPONSE)
232+
https_server.serve("/inventory", b"not sent for HEAD" if method == "HEAD" else b"", status=status)
233+
subject = client(https_server)
234+
235+
response = subject.request(method, https_server.url + "/inventory")
236+
237+
assert response.status == status
238+
assert response.data == b""
239+
assert response.data.decode() == ""
240+
assert [request[1] for request in https_server.requests] == ["/token", "/inventory"]
241+
242+
243+
@pytest.mark.parametrize("endpoint", ["token", "inventory"])
244+
def test_malformed_response_headers_fail_without_logging_credentials(https_server, caplog, endpoint):
245+
private_data = "Bearer local-test-private-token"
246+
https_server.serve("/token", TOKEN_RESPONSE)
247+
https_server.serve("/inventory", {"items": [123]})
248+
payload = TOKEN_RESPONSE if endpoint == "token" else {"items": [123]}
249+
https_server.serve(f"/{endpoint}", payload, headers={"Broken header": private_data})
250+
subject = client(https_server)
251+
expected = TokenExchangeError if endpoint == "token" else DownstreamRequestError
252+
253+
with pytest.raises(expected) as error:
254+
subject.request("GET", https_server.url + "/inventory")
255+
256+
assert not error.value.retryable
257+
assert error.value.__context__ is None
258+
assert error.value.__cause__ is None
259+
assert private_data not in "".join(traceback.format_exception(error.value))
260+
assert private_data not in caplog.text
261+
assert not [record for record in caplog.records if record.name == "urllib3.connection"]
262+
expected_paths = ["/token"] if endpoint == "token" else ["/token", "/inventory"]
263+
assert [request[1] for request in https_server.requests] == expected_paths
264+
265+
https_server.serve(f"/{endpoint}", payload)
266+
assert subject.request("GET", https_server.url + "/inventory").json() == {"items": [123]}
267+
268+
269+
def test_valid_extended_header_values_are_sent_over_tls(https_server):
270+
https_server.serve("/token", TOKEN_RESPONSE)
271+
https_server.serve("/inventory", {"items": [123]})
272+
value = "caf\xe9\t\x80\xff"
273+
274+
response = client(https_server).request(
275+
"GET",
276+
https_server.url + "/inventory",
277+
headers={"X-Trace": value},
278+
)
279+
280+
assert response.status == 200
281+
assert https_server.requests[-1][2]["X-Trace"] == value

0 commit comments

Comments
 (0)