diff --git a/aws_lambda_powertools/utilities/auth_alpha/__init__.py b/aws_lambda_powertools/utilities/auth_alpha/__init__.py index 5809a0dec7e..2090e16bd89 100644 --- a/aws_lambda_powertools/utilities/auth_alpha/__init__.py +++ b/aws_lambda_powertools/utilities/auth_alpha/__init__.py @@ -6,15 +6,21 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: + from aws_lambda_powertools.utilities.auth_alpha.exceptions import AuthFailureReason as AuthFailureReason from aws_lambda_powertools.utilities.auth_alpha.jwt import AuthErrorContext as AuthErrorContext - from aws_lambda_powertools.utilities.auth_alpha.jwt import AuthFailureReason as AuthFailureReason from aws_lambda_powertools.utilities.auth_alpha.jwt import JWTVerifier as JWTVerifier + from aws_lambda_powertools.utilities.auth_alpha.oauth2 import OAuth2Client as OAuth2Client -__all__ = ["AuthErrorContext", "AuthFailureReason", "JWTVerifier"] +__all__ = ["AuthErrorContext", "AuthFailureReason", "JWTVerifier", "OAuth2Client"] def __getattr__(name: str) -> object: - modules = {"AuthErrorContext": "jwt", "AuthFailureReason": "jwt", "JWTVerifier": "jwt"} + modules = { + "AuthErrorContext": "jwt", + "AuthFailureReason": "exceptions", + "JWTVerifier": "jwt", + "OAuth2Client": "oauth2", + } if name in modules: value = getattr(importlib.import_module(f"{__name__}.{modules[name]}"), name) globals()[name] = value diff --git a/aws_lambda_powertools/utilities/auth_alpha/_internal/errors.py b/aws_lambda_powertools/utilities/auth_alpha/_internal/errors.py new file mode 100644 index 00000000000..89e0f2de078 --- /dev/null +++ b/aws_lambda_powertools/utilities/auth_alpha/_internal/errors.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from functools import wraps +from typing import TYPE_CHECKING, ParamSpec, TypeVar + +from aws_lambda_powertools.utilities.auth_alpha.exceptions import AuthError + +if TYPE_CHECKING: + from collections.abc import Callable + +_P = ParamSpec("_P") +_T = TypeVar("_T") + + +def sanitize_errors(operation: Callable[_P, _T]) -> Callable[_P, _T]: + """Detach provider exceptions before an Auth error leaves a public operation.""" + + @wraps(operation) + def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _T: + try: + return operation(*args, **kwargs) + except AuthError as error: + # `raise ... from None` only suppresses display of the context. + # Clear both references and use a bare re-raise so Python does not + # attach the active exception again. + error.__context__ = None + error.__cause__ = None + raise + + return wrapper diff --git a/aws_lambda_powertools/utilities/auth_alpha/_internal/http.py b/aws_lambda_powertools/utilities/auth_alpha/_internal/http.py index 109ebca4836..b7577d24ed1 100644 --- a/aws_lambda_powertools/utilities/auth_alpha/_internal/http.py +++ b/aws_lambda_powertools/utilities/auth_alpha/_internal/http.py @@ -1,16 +1,23 @@ from __future__ import annotations import json -from typing import TYPE_CHECKING, Any +from io import BytesIO +from typing import TYPE_CHECKING, Any, cast import urllib3 -from urllib3.connection import HTTPConnection +from urllib3.poolmanager import pool_classes_by_scheme from aws_lambda_powertools.utilities.auth_alpha._internal.deadline import Deadline, RequestError +from aws_lambda_powertools.utilities.auth_alpha._internal.transport import ( + DeadlineHTTPSConnectionPool, + response_deadline, +) if TYPE_CHECKING: from collections.abc import Mapping + from urllib3.connectionpool import HTTPConnectionPool + _MAX_JSON_BYTES = 1024 * 1024 @@ -19,6 +26,27 @@ class HTTPClient: def __init__(self) -> None: self.pool = urllib3.PoolManager(cert_reqs="CERT_REQUIRED") + self.pool.pool_classes_by_scheme = pool_classes_by_scheme.copy() + pool_classes = cast("dict[str, type[HTTPConnectionPool]]", self.pool.pool_classes_by_scheme) + pool_classes["https"] = DeadlineHTTPSConnectionPool + + def _open_response( + self, + method: str, + url: str, + deadline: Deadline, + **options: Any, + ) -> urllib3.response.BaseHTTPResponse: + with response_deadline(deadline): + return self.pool.request( + method, + url, + timeout=urllib3.Timeout(total=deadline.remaining()), + retries=False, + redirect=False, + preload_content=False, + **options, + ) def json_request( self, @@ -31,15 +59,12 @@ def json_request( ) -> tuple[int, dict[str, Any]]: response = None try: - response = self.pool.request( + response = self._open_response( method, url, + deadline, body=body, headers=headers, - timeout=urllib3.Timeout(total=deadline.remaining()), - retries=False, - redirect=False, - preload_content=False, ) if response.status != 200: deadline.remaining() @@ -55,23 +80,67 @@ def json_request( @staticmethod def _read_json(response: urllib3.response.BaseHTTPResponse, deadline: Deadline) -> dict[str, Any]: + content = HTTPClient._read_body(response, deadline, limit=_MAX_JSON_BYTES, decode_content=False) + try: + data = json.loads(content) + except (ValueError, UnicodeError, RecursionError): + raise RequestError() from None + if not isinstance(data, dict): + raise RequestError() + return data + + def request( + self, + method: str, + url: str, + deadline: Deadline, + *, + headers: Mapping[str, str], + **options: Any, + ) -> urllib3.response.HTTPResponse: + """Buffer an authenticated response within one network time budget.""" + response = None + try: + response = self._open_response( + method, + url, + deadline, + headers=headers, + **options, + ) + content = self._read_body(response, deadline) + return urllib3.HTTPResponse( + body=BytesIO(content), + status=response.status, + headers=response.headers, + reason=response.reason, + version=response.version, + request_method=method, + request_url=url, + decode_content=False, + ) + finally: + if response is not None: + response.close() + response.release_conn() + + @staticmethod + def _read_body( + response: urllib3.response.BaseHTTPResponse, + deadline: Deadline, + *, + limit: int | None = None, + decode_content: bool = True, + ) -> bytes: chunks = bytearray() while True: - remaining = deadline.remaining() - connection = response.connection - if isinstance(connection, HTTPConnection) and connection.sock is not None: - connection.sock.settimeout(remaining) - chunk = response.read1(min(65536, _MAX_JSON_BYTES + 1 - len(chunks)), decode_content=False) + deadline.remaining() + size = 65536 if limit is None else min(65536, limit + 1 - len(chunks)) + chunk = response.read1(size, decode_content=decode_content) deadline.remaining() if not chunk: break chunks.extend(chunk) - if len(chunks) > _MAX_JSON_BYTES: + if limit is not None and len(chunks) > limit: raise RequestError() - try: - data = json.loads(chunks) - except (ValueError, UnicodeError, RecursionError): - raise RequestError() from None - if not isinstance(data, dict): - raise RequestError() - return data + return bytes(chunks) diff --git a/aws_lambda_powertools/utilities/auth_alpha/_internal/scopes.py b/aws_lambda_powertools/utilities/auth_alpha/_internal/scopes.py new file mode 100644 index 00000000000..d80e0901f29 --- /dev/null +++ b/aws_lambda_powertools/utilities/auth_alpha/_internal/scopes.py @@ -0,0 +1,14 @@ +from __future__ import annotations + +from aws_lambda_powertools.utilities.auth_alpha._internal.validation import string_list + + +def valid_scope(value: str) -> bool: + return bool(value) and all(33 <= ord(character) <= 126 and character not in {'"', "\\"} for character in value) + + +def required_scopes(scopes: list[str] | None) -> tuple[str, ...]: + values = string_list(scopes if scopes is not None else []) + if not all(valid_scope(value) for value in values): + raise ValueError("Scopes must be valid OAuth scope tokens") + return values diff --git a/aws_lambda_powertools/utilities/auth_alpha/_internal/transport.py b/aws_lambda_powertools/utilities/auth_alpha/_internal/transport.py new file mode 100644 index 00000000000..6620025eb62 --- /dev/null +++ b/aws_lambda_powertools/utilities/auth_alpha/_internal/transport.py @@ -0,0 +1,101 @@ +from __future__ import annotations + +from contextlib import contextmanager +from contextvars import ContextVar +from http.client import HTTPResponse +from io import BufferedReader, RawIOBase +from typing import TYPE_CHECKING + +from urllib3.connection import HTTPSConnection +from urllib3.connectionpool import HTTPSConnectionPool +from urllib3.exceptions import HeaderParsingError +from urllib3.util.response import assert_header_parsing + +from aws_lambda_powertools.utilities.auth_alpha._internal.deadline import RequestError + +if TYPE_CHECKING: + from collections.abc import Iterator + from socket import socket + from typing import Protocol + + from typing_extensions import Buffer + + from aws_lambda_powertools.utilities.auth_alpha._internal.deadline import Deadline + + class _ResponseStream(Protocol): + def readinto1(self, buffer: Buffer, /) -> int: ... + def close(self) -> None: ... + + +_current_deadline: ContextVar[Deadline] = ContextVar("auth_response_deadline") + + +@contextmanager +def response_deadline(deadline: Deadline) -> Iterator[None]: + """Pass the operation's deadline to responses created by this synchronous call.""" + token = _current_deadline.set(deadline) + try: + yield + finally: + _current_deadline.reset(token) + + +class _DeadlineReader(RawIOBase): + """Check the original budget on every refill, including HTTP framing reads.""" + + def __init__(self, stream: _ResponseStream, sock: socket, deadline: Deadline) -> None: + super().__init__() + self._stream = stream + self._socket = sock + self._deadline = deadline + + def readable(self) -> bool: + return True + + def readinto(self, buffer: Buffer, /) -> int: + self._socket.settimeout(self._deadline.remaining()) + # Unlike readinto(), readinto1() performs at most one underlying read. + count = self._stream.readinto1(buffer) + self._deadline.remaining() + return count + + def close(self) -> None: + try: + self._stream.close() + finally: + super().close() + + +class _DeadlineResponse(HTTPResponse): + def __init__( + self, + sock: socket, + debuglevel: int = 0, + method: str | None = None, + url: str | None = None, + ) -> None: + deadline = _current_deadline.get() + super().__init__(sock, debuglevel=debuglevel, method=method, url=url) + # Retain the wrapper through body consumption: read1() can also parse + # chunk-size lines, delimiters and trailers before returning to our loop. + # The stream owns the socket reference even for Connection: close. + self.fp = BufferedReader(_DeadlineReader(self.fp, sock, deadline), buffer_size=8192) + + def begin(self) -> None: + super().begin() + try: + # urllib3 otherwise logs malformed headers, including provider data, + # before the public Auth operation can sanitize the failure. + assert_header_parsing(self.msg) + except (HeaderParsingError, TypeError): + raise RequestError() from None + + +class _DeadlineHTTPSConnection(HTTPSConnection): + response_class = _DeadlineResponse + + +class DeadlineHTTPSConnectionPool(HTTPSConnectionPool): + """Use the stdlib response hook without overriding urllib3's request machinery.""" + + ConnectionCls = _DeadlineHTTPSConnection diff --git a/aws_lambda_powertools/utilities/auth_alpha/exceptions.py b/aws_lambda_powertools/utilities/auth_alpha/exceptions.py new file mode 100644 index 00000000000..e6f598183e4 --- /dev/null +++ b/aws_lambda_powertools/utilities/auth_alpha/exceptions.py @@ -0,0 +1,29 @@ +"""Credential-free errors raised by the Auth utility.""" + +from enum import Enum + + +class AuthFailureReason(str, Enum): + """Stable, credential-free reasons suitable for application logs and metrics.""" + + MISSING_TOKEN = "missing_token" # nosec B105 + INVALID_TOKEN = "invalid_token" # nosec B105 + INVALID_CLAIMS = "invalid_claims" + TOKEN_EXPIRED = "token_expired" # nosec B105 + INVALID_SIGNATURE = "invalid_signature" + INSUFFICIENT_SCOPE = "insufficient_scope" + FORBIDDEN = "forbidden" + JWKS_UNAVAILABLE = "jwks_unavailable" + TOKEN_EXCHANGE_FAILED = "token_exchange_failed" # nosec B105 + DOWNSTREAM_REQUEST_FAILED = "downstream_request_failed" + + +class AuthError(Exception): + """Base error with a fixed message that never includes credential material.""" + + message = "Authentication failed" + reason = AuthFailureReason.INVALID_TOKEN + retryable = False + + def __init__(self) -> None: + super().__init__(self.message) diff --git a/aws_lambda_powertools/utilities/auth_alpha/jwt/_internal/authorization.py b/aws_lambda_powertools/utilities/auth_alpha/jwt/_internal/authorization.py index b035991901e..6732104425e 100644 --- a/aws_lambda_powertools/utilities/auth_alpha/jwt/_internal/authorization.py +++ b/aws_lambda_powertools/utilities/auth_alpha/jwt/_internal/authorization.py @@ -3,7 +3,7 @@ from collections.abc import Mapping from typing import Any -from aws_lambda_powertools.utilities.auth_alpha._internal.validation import string_list +from aws_lambda_powertools.utilities.auth_alpha._internal.scopes import required_scopes, valid_scope from aws_lambda_powertools.utilities.auth_alpha.jwt.exceptions import ( AuthError, AuthFailureReason, @@ -11,6 +11,8 @@ InvalidTokenError, ) +__all__ = ["required_scopes"] + class MissingTokenError(InvalidTokenError): """No authorization header was supplied.""" @@ -65,17 +67,6 @@ def _authorization_values(headers: Any) -> list[Any]: return values -def valid_scope(value: str) -> bool: - return bool(value) and all(33 <= ord(character) <= 126 and character not in {'"', "\\"} for character in value) - - -def required_scopes(scopes: list[str] | None) -> tuple[str, ...]: - values = string_list(scopes if scopes is not None else []) - if not all(valid_scope(value) for value in values): - raise ValueError("Scopes must be valid OAuth scope tokens") - return values - - def enforce_scopes(claims: dict[str, Any], expected: tuple[str, ...]) -> None: value: Any = next((claims[name] for name in ("scope", "scp", "scopes") if name in claims), []) if isinstance(value, str): diff --git a/aws_lambda_powertools/utilities/auth_alpha/jwt/_internal/errors.py b/aws_lambda_powertools/utilities/auth_alpha/jwt/_internal/errors.py index c99f88e54d2..0ac7598fe60 100644 --- a/aws_lambda_powertools/utilities/auth_alpha/jwt/_internal/errors.py +++ b/aws_lambda_powertools/utilities/auth_alpha/jwt/_internal/errors.py @@ -1,30 +1,3 @@ -from __future__ import annotations +from aws_lambda_powertools.utilities.auth_alpha._internal.errors import sanitize_errors -from functools import wraps -from typing import TYPE_CHECKING, ParamSpec, TypeVar - -from aws_lambda_powertools.utilities.auth_alpha.jwt.exceptions import AuthError - -if TYPE_CHECKING: - from collections.abc import Callable - -_P = ParamSpec("_P") -_T = TypeVar("_T") - - -def sanitize_errors(operation: Callable[_P, _T]) -> Callable[_P, _T]: - """Detach provider exceptions before an Auth error leaves a public operation.""" - - @wraps(operation) - def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _T: - try: - return operation(*args, **kwargs) - except AuthError as error: - # `raise ... from None` only suppresses display of the context. - # Clear both references and use a bare re-raise so Python does not - # attach the active exception again. - error.__context__ = None - error.__cause__ = None - raise - - return wrapper +__all__ = ["sanitize_errors"] diff --git a/aws_lambda_powertools/utilities/auth_alpha/jwt/exceptions.py b/aws_lambda_powertools/utilities/auth_alpha/jwt/exceptions.py index 6c460b34790..03f2044f302 100644 --- a/aws_lambda_powertools/utilities/auth_alpha/jwt/exceptions.py +++ b/aws_lambda_powertools/utilities/auth_alpha/jwt/exceptions.py @@ -1,30 +1,16 @@ -"""Credential-free errors raised by the Auth utility.""" - -from enum import Enum - - -class AuthFailureReason(str, Enum): - """Stable, credential-free reasons suitable for application logs and metrics.""" - - MISSING_TOKEN = "missing_token" # nosec B105 - INVALID_TOKEN = "invalid_token" # nosec B105 - INVALID_CLAIMS = "invalid_claims" - TOKEN_EXPIRED = "token_expired" # nosec B105 - INVALID_SIGNATURE = "invalid_signature" - INSUFFICIENT_SCOPE = "insufficient_scope" - FORBIDDEN = "forbidden" - JWKS_UNAVAILABLE = "jwks_unavailable" - - -class AuthError(Exception): - """Base error with a fixed message that never includes credential material.""" - - message = "Authentication failed" - reason = AuthFailureReason.INVALID_TOKEN - retryable = False - - def __init__(self) -> None: - super().__init__(self.message) +"""Credential-free errors raised by JWT verification.""" + +from aws_lambda_powertools.utilities.auth_alpha.exceptions import AuthError, AuthFailureReason + +__all__ = [ + "AuthError", + "AuthFailureReason", + "InvalidTokenError", + "InvalidClaimsError", + "TokenExpiredError", + "InvalidSignatureError", + "JWKSFetchError", +] class InvalidTokenError(AuthError): diff --git a/aws_lambda_powertools/utilities/auth_alpha/oauth2/__init__.py b/aws_lambda_powertools/utilities/auth_alpha/oauth2/__init__.py new file mode 100644 index 00000000000..51c2bfea71e --- /dev/null +++ b/aws_lambda_powertools/utilities/auth_alpha/oauth2/__init__.py @@ -0,0 +1,23 @@ +"""OAuth 2.0 client-credentials token acquisition.""" + +from __future__ import annotations + +import importlib +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from aws_lambda_powertools.utilities.auth_alpha.oauth2.client import OAuth2Client as OAuth2Client + +__all__ = ["OAuth2Client"] + + +def __getattr__(name: str) -> object: + if name == "OAuth2Client": + value = importlib.import_module(f"{__name__}.client").OAuth2Client + globals()[name] = value + return value + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + +def __dir__() -> list[str]: + return sorted(set(globals()) | set(__all__)) diff --git a/aws_lambda_powertools/utilities/auth_alpha/oauth2/client.py b/aws_lambda_powertools/utilities/auth_alpha/oauth2/client.py new file mode 100644 index 00000000000..acb49b8a8bc --- /dev/null +++ b/aws_lambda_powertools/utilities/auth_alpha/oauth2/client.py @@ -0,0 +1,387 @@ +from __future__ import annotations + +import base64 +import re +import threading +import time +from collections.abc import Mapping +from dataclasses import dataclass, field +from string import ascii_letters, digits, hexdigits +from typing import TYPE_CHECKING, Any, Literal +from urllib.parse import quote_plus, urlencode, urlsplit + +import urllib3 + +from aws_lambda_powertools.utilities.auth_alpha._internal.deadline import Deadline, RequestError +from aws_lambda_powertools.utilities.auth_alpha._internal.errors import sanitize_errors +from aws_lambda_powertools.utilities.auth_alpha._internal.http import HTTPClient +from aws_lambda_powertools.utilities.auth_alpha._internal.scopes import required_scopes +from aws_lambda_powertools.utilities.auth_alpha._internal.validation import ( + finite_seconds, + https_url, + is_nonempty_string, +) +from aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions import DownstreamRequestError, TokenExchangeError + +if TYPE_CHECKING: + from collections.abc import Callable + +_BEARER_TOKEN = re.compile(r"[-A-Za-z0-9._~+/]+=*") +_HEADER_NAME = re.compile(r"[-!#$%&'*+.^_`|~0-9A-Za-z]+") +# HTTP field values allow horizontal tabs, visible ASCII, and extended Latin-1 bytes. +_HEADER_VALUE = re.compile(r"[\t\x20-\x7e\x80-\xff]*") +_RESOURCE_CHARACTERS = frozenset(ascii_letters + digits + "-._~:/?[]@!$&'()*+,;=") +_HEX_DIGITS = frozenset(hexdigits) + + +def _valid_resource_characters(resource: str) -> bool: + """Scan once, accepting URI characters and complete percent escapes, but no fragment.""" + characters = iter(resource) + for character in characters: + if character == "%": + if next(characters, "") not in _HEX_DIGITS or next(characters, "") not in _HEX_DIGITS: + return False + elif character not in _RESOURCE_CHARACTERS: + return False + return True + + +@dataclass(frozen=True) +class _AccessToken: + value: str = field(repr=False) + expires_at: float | None + + def cacheable(self) -> bool: + return self.expires_at is not None and time.monotonic() < self.expires_at - 30 + + def usable(self) -> bool: + return self.expires_at is None or time.monotonic() < self.expires_at + + +@dataclass +class _Exchange: + done: threading.Event = field(default_factory=threading.Event, repr=False) + token: _AccessToken | None = field(default=None, repr=False) + retryable: bool = False + + +class OAuth2Client: + """Acquire bearer tokens using client credentials for one configured resource. + + Parameters + ---------- + token_url : str + Trusted HTTPS OAuth token endpoint. + client_id : str + OAuth client identifier. + client_secret : str | Callable[[], str] + Secret or loader invoked for each exchange attempt. + auth_method : Literal["client_secret_basic", "client_secret_post"] + Client authentication method, by default ``client_secret_basic``. + ``client_secret_post`` sends credentials in the form body without an + Authorization header. The client never switches methods automatically. + scopes : list[str], optional + Scopes requested on every exchange. + audience : str, optional + Provider-specific audience request field, mutually exclusive with resource. + resource : str, optional + RFC 8707 resource request field, mutually exclusive with audience. + Must be an absolute URI without a fragment; URNs are supported. + timeout_seconds : float + Positive acquisition budget including retries, by default 3. + + Notes + ----- + Instances do not share tokens. Tokens are reacquired 30 seconds before + expiration. Short-lived tokens and tokens without a lifetime are not cached. + Configure timeouts on application-provided secret loaders. + + Examples + -------- + ```python + client = OAuth2Client( + token_url="https://idp.example.com/token", + client_id="orders", + client_secret=load_secret, + resource="https://inventory.example.com", + scopes=["inventory:read"], + ) + headers = client.auth_headers() + ``` + """ + + def __init__( + self, + *, + token_url: str, + client_id: str, + client_secret: str | Callable[[], str], + auth_method: Literal["client_secret_basic", "client_secret_post"] = "client_secret_basic", + scopes: list[str] | None = None, + audience: str | None = None, + resource: str | None = None, + timeout_seconds: float = 3, + ) -> None: + self._token_url = https_url(token_url) + if not is_nonempty_string(client_id): + raise ValueError("A nonempty OAuth client ID is required") + if not callable(client_secret) and (not isinstance(client_secret, str) or not client_secret): + raise ValueError("client_secret must be a nonempty string or a callable") + if auth_method not in ("client_secret_basic", "client_secret_post"): + raise ValueError("auth_method must be client_secret_basic or client_secret_post") + if audience is not None and resource is not None: + raise ValueError("audience and resource are mutually exclusive") + self._client_id = client_id + self._client_secret = client_secret + self._auth_method = auth_method + self._scopes = required_scopes(scopes) + self._timeout = finite_seconds(timeout_seconds, positive=True) + self._fields = {"grant_type": "client_credentials"} + if self._scopes: + self._fields["scope"] = " ".join(self._scopes) + for name, value in (("audience", audience), ("resource", resource)): + if value is not None: + if not is_nonempty_string(value): + raise ValueError("Resource selection must be a nonempty string") + self._fields[name] = value + try: + self._client_id.encode("utf-8") + self._body = urlencode(self._fields).encode() + except UnicodeError: + raise ValueError("OAuth configuration must contain valid UTF-8 strings") from None + if resource is not None: + self._validate_resource(resource) + self._http = HTTPClient() + self._cached_token: _AccessToken | None = None + self._flight: _Exchange | None = None + self._lock = threading.Lock() + + def __repr__(self) -> str: + return "" + + @staticmethod + def _validate_resource(resource: str) -> None: + try: + scheme = urlsplit(resource).scheme + # Older urlsplit versions also recognize schemes starting with a digit. + valid = scheme[:1].isalpha() and _valid_resource_characters(resource) + except ValueError: + valid = False + if not valid: + raise ValueError("resource must be an absolute URI without a fragment") + + @sanitize_errors + def auth_headers(self) -> dict[str, str]: + """Return an Authorization header for this client's configured resource. + + Raises + ------ + TokenExchangeError + A usable bearer token could not be obtained within the budget. + + Examples + -------- + ```python + headers = client.auth_headers() + response = http.request("GET", trusted_inventory_url, headers=headers) + ``` + """ + try: + token = self._get_token(Deadline(self._timeout)) + except RequestError as error: + raise TokenExchangeError(retryable=error.retryable) from None + return {"Authorization": f"Bearer {token.value}"} + + @sanitize_errors + def request( + self, + method: str, + url: str, + *, + timeout: float = 5, + headers: Mapping[str, str] | None = None, + **options: Any, + ) -> urllib3.response.BaseHTTPResponse: + """Send a synchronous HTTPS request using this resource's bearer token. + + Only trusted destination URLs should be supplied. Redirects and retries + are disabled, and an existing Authorization header is rejected. + ``body``, ``fields``, ``json``, ``encode_multipart`` and + ``multipart_boundary`` are forwarded to urllib3. + + Parameters + ---------- + method : str + HTTP method. + url : str + Trusted HTTPS destination for this resource's credentials. + timeout : float + Positive downstream timeout, separate from acquisition, by default 5. + headers : Mapping[str, str], optional + Additional headers with HTTP token names, excluding Authorization. + + Returns + ------- + urllib3.response.BaseHTTPResponse + Downstream response; inspect its status before consuming its body. + + Raises + ------ + TokenExchangeError + Token acquisition failed. + DownstreamRequestError + Downstream transport failed; the request is never replayed automatically. + ValueError + Request configuration is invalid. + + Examples + -------- + ```python + response = client.request("GET", "https://inventory.example.com/items") + if response.status == 200: + items = response.json() + ``` + """ + target = https_url(url) + duration = finite_seconds(timeout, positive=True) + allowed = {"body", "fields", "json", "encode_multipart", "multipart_boundary"} + if not options.keys() <= allowed: + raise ValueError("Unsupported authenticated request option") + if not isinstance(method, str) or not re.fullmatch(r"[A-Za-z]+", method): + raise ValueError("A valid HTTP method is required") + request_headers = self._request_headers(headers) + request_headers.update(self.auth_headers()) + deadline = Deadline(duration) + try: + return self._http.request( + method.upper(), + target, + deadline, + headers=request_headers, + **options, + ) + except (urllib3.exceptions.HTTPError, OSError, ValueError, TypeError, RequestError): + raise DownstreamRequestError() from None + + @staticmethod + def _request_headers(headers: Mapping[str, str] | None) -> dict[str, str]: + if headers is None: + return {} + if not isinstance(headers, Mapping): + raise ValueError("Request headers must be a mapping of strings") + for name, value in headers.items(): + if ( + not isinstance(name, str) + or not isinstance(value, str) + or not _HEADER_NAME.fullmatch(name) + or name.lower() == "authorization" + or not _HEADER_VALUE.fullmatch(value) + ): + raise ValueError("Request headers must be valid and must not include Authorization") + return dict(headers) + + def _get_token(self, deadline: Deadline) -> _AccessToken: + with self._lock: + if self._cached_token is not None and self._cached_token.cacheable(): + return self._cached_token + self._cached_token = None + owner = self._flight is None + if self._flight is None: + self._flight = _Exchange() + flight = self._flight + if owner: + self._run_exchange(flight, deadline) + elif not flight.done.wait(timeout=deadline.remaining()): + raise TokenExchangeError(retryable=True) + deadline.remaining() + if flight.token is None: + raise TokenExchangeError(retryable=flight.retryable) + if not flight.token.usable(): + raise TokenExchangeError() + return flight.token + + def _run_exchange(self, flight: _Exchange, deadline: Deadline) -> None: + try: + token = self._exchange(deadline) + with self._lock: + if token.cacheable(): + self._cached_token = token + flight.token = token + except (TokenExchangeError, RequestError) as error: + flight.retryable = error.retryable + raise + finally: + # Waiters keep this flight's result, including uncacheable short + # tokens. Calls starting after completion must acquire their own. + with self._lock: + self._flight = None + flight.done.set() + + def _exchange(self, deadline: Deadline) -> _AccessToken: + retry_delays = iter((0.1, 0.2)) + while True: + try: + return self._exchange_once(deadline) + except RequestError as error: + delay = next(retry_delays, None) + if not error.retryable or delay is None or deadline.remaining() <= delay: + raise TokenExchangeError(retryable=error.retryable) from None + time.sleep(delay) + + def _token_request(self) -> tuple[bytes, dict[str, str]]: + try: + secret = self._client_secret if isinstance(self._client_secret, str) else self._client_secret() + except Exception: + # Secret providers can raise arbitrary exceptions containing their + # configuration or response data. None of it crosses this boundary. + raise TokenExchangeError() from None + if not isinstance(secret, str) or not secret: + raise TokenExchangeError() + headers = {"Content-Type": "application/x-www-form-urlencoded"} + try: + if self._auth_method == "client_secret_post": + # Keep credentials out of reusable fields and load them for every attempt. + fields = {**self._fields, "client_id": self._client_id, "client_secret": secret} + return urlencode(fields).encode(), headers + credentials = f"{quote_plus(self._client_id)}:{quote_plus(secret)}" + except UnicodeError: + raise TokenExchangeError() from None + headers["Authorization"] = f"Basic {base64.b64encode(credentials.encode()).decode()}" + return self._body, headers + + def _exchange_once(self, deadline: Deadline) -> _AccessToken: + body, headers = self._token_request() + started = time.monotonic() + status, payload = self._http.json_request( + "POST", + self._token_url, + deadline, + body=body, + headers=headers, + ) + if status != 200: + raise RequestError(retryable=status == 429 or 500 <= status <= 599) + return self._parse_token(payload, started) + + @staticmethod + def _parse_token(payload: dict[str, Any], started: float) -> _AccessToken: + value = payload.get("access_token") + token_type = payload.get("token_type") + if ( + not isinstance(value, str) + or not _BEARER_TOKEN.fullmatch(value) + or not isinstance(token_type, str) + or token_type.lower() != "bearer" + ): + raise TokenExchangeError() + expires_at = None + if "expires_in" in payload: + try: + lifetime = finite_seconds(payload["expires_in"], positive=True) + except ValueError: + raise TokenExchangeError() from None + expires_at = started + lifetime + token = _AccessToken(value, expires_at) + if not token.usable(): + raise TokenExchangeError() + return token diff --git a/aws_lambda_powertools/utilities/auth_alpha/oauth2/exceptions.py b/aws_lambda_powertools/utilities/auth_alpha/oauth2/exceptions.py new file mode 100644 index 00000000000..f4cdd7f060e --- /dev/null +++ b/aws_lambda_powertools/utilities/auth_alpha/oauth2/exceptions.py @@ -0,0 +1,32 @@ +"""Credential-free errors raised by outbound OAuth requests.""" + +from aws_lambda_powertools.utilities.auth_alpha.exceptions import AuthError, AuthFailureReason + +__all__ = ["AuthError", "AuthFailureReason", "TokenExchangeError", "DownstreamRequestError"] + + +class TokenExchangeError(AuthError): + """A usable bearer token could not be obtained. + + ``retryable`` indicates a transient endpoint failure or acquisition timeout. + Invalid responses, rejected credentials, and secret-loader failures are not + retried automatically. Exception messages never include provider details. + """ + + message = "Unable to acquire an access token" + reason = AuthFailureReason.TOKEN_EXCHANGE_FAILED + + def __init__(self, *, retryable: bool = False) -> None: + self.retryable = retryable + super().__init__() + + +class DownstreamRequestError(AuthError): + """The authenticated HTTP operation could not complete. + + The server may already have performed the operation. Callers must decide + whether replay is safe; this error does not advise automatic retries. + """ + + message = "Authenticated request failed" + reason = AuthFailureReason.DOWNSTREAM_REQUEST_FAILED diff --git a/docs/api_doc/auth_alpha.md b/docs/api_doc/auth_alpha.md index 52409094848..30854875ab5 100644 --- a/docs/api_doc/auth_alpha.md +++ b/docs/api_doc/auth_alpha.md @@ -5,3 +5,5 @@ ::: aws_lambda_powertools.utilities.auth_alpha.AuthErrorContext ::: aws_lambda_powertools.utilities.auth_alpha.jwt.exceptions ::: aws_lambda_powertools.utilities.auth_alpha.jwt.testing +::: aws_lambda_powertools.utilities.auth_alpha.oauth2.client +::: aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions diff --git a/docs/getting-started/install.md b/docs/getting-started/install.md index e744b9e35e0..f17a68ee92b 100644 --- a/docs/getting-started/install.md +++ b/docs/getting-started/install.md @@ -43,6 +43,7 @@ Some features require additional dependencies. Install them as needed: | [Validation](../utilities/validation.md) | `pip install "aws-lambda-powertools[validation]"` | `fastjsonschema` | | [Parser](../utilities/parser.md) | `pip install "aws-lambda-powertools[parser]"` | `pydantic` | | [JWT verification (alpha)](../utilities/auth.md) | `pip install "aws-lambda-powertools[jwt]"` | `PyJWT`, `cryptography`, `urllib3` | +| [OAuth2 client (alpha)](../utilities/oauth2.md) | `pip install "aws-lambda-powertools[oauth2]"` | `urllib3` | | [Data Masking](../utilities/data_masking.md) | `pip install "aws-lambda-powertools[datamasking]"` | `aws-encryption-sdk`, `jsonpath-ng` | | [Datadog Metrics](../core/metrics/datadog.md) | `pip install "aws-lambda-powertools[datadog]"` | `datadog-lambda` | | [Kafka (Avro)](../utilities/kafka.md) | `pip install "aws-lambda-powertools[kafka-consumer-avro]"` | `avro` | diff --git a/docs/index.md b/docs/index.md index 475eda10a74..6a4c1317b0f 100644 --- a/docs/index.md +++ b/docs/index.md @@ -54,7 +54,7 @@ Powertools for AWS Lambda (Python) is a developer toolkit to implement Serverles | [Metrics](./core/metrics.md) | Custom Metrics created asynchronously via CloudWatch Embedded Metric Format (EMF) | | [Event Handler](./core/event_handler/api_gateway.md) | Event handler for API Gateway, ALB, Lambda Function URL, VPC Lattice, AppSync, and Bedrock Agents | | [Parameters](./utilities/parameters.md) | Retrieve and cache parameter values from Parameter Store, Secrets Manager, AppConfig, or DynamoDB | -| [Auth (alpha)](./utilities/auth.md) | Verify JWT access tokens and use verified claims in Lambda workloads | +| [Auth (alpha)](./utilities/auth.md) | Verify JWT access tokens and [acquire client-credentials tokens](./utilities/oauth2.md) for downstream APIs | | [Parser](./utilities/parser.md) | Data parsing and deep validation using Pydantic | | [Batch Processing](./utilities/batch.md) | Handle partial failures for SQS, Kinesis Data Streams, and DynamoDB Streams | | [Idempotency](./utilities/idempotency.md) | Make your Lambda functions idempotent and prevent duplicate execution | diff --git a/docs/utilities/auth.md b/docs/utilities/auth.md index 849f4ef3abd..12c5aced9aa 100644 --- a/docs/utilities/auth.md +++ b/docs/utilities/auth.md @@ -9,6 +9,8 @@ status: new Auth verifies JWT access tokens in any Lambda workload. Use `verify()` directly, create Event Handler middleware with `require()`, or build a Lambda authorizer response with `authorize()`. +For client-credentials tokens used to call a downstream API, see [OAuth2 client (alpha)](oauth2.md). + ```mermaid flowchart LR Token["JWT access token"] --> Verify["JWTVerifier.verify()"] @@ -36,6 +38,20 @@ pip install "aws-lambda-powertools[jwt]" The `jwt` extra installs PyJWT, cryptography, and urllib3. Build dependencies for the same Python version and architecture as your Lambda function. See [cross-platform builds](../build_recipes/cross-platform.md). +#### Compatibility with boto3 + +If the function also uses boto3, including through Parameters, install and validate the SDK together with JWT verification: + +```shell +pip install "aws-lambda-powertools[jwt,aws-sdk]" +python -m pip check +``` + +Resolve all function dependencies together, lock their versions, and package boto3, botocore, and urllib3 together. Make a failed `pip check` fail the build. + +Packaged urllib3 overrides the runtime copy and can conflict with the runtime's SDK. +Build checks cannot validate dependencies supplied only by the runtime or separate layers. See [AWS packaging guidance](https://docs.aws.amazon.com/lambda/latest/dg/python-package.html#python-package-dependencies). + ### Create a verifier Create the verifier outside the Lambda handler so warm invocations reuse its signing-key cache. Configure the trusted issuer, this workload's audience, and the algorithms accepted from that issuer. @@ -62,11 +78,12 @@ Replace `token_use` with the access-token marker used by your identity provider. ### Required resources -JWT verification requires no additional IAM permissions. When using issuer discovery or a remote JWKS endpoint, the function needs outbound HTTPS access to the identity provider. A function in private subnets might need a NAT gateway or private connectivity. Static `jwks` does not use the network, but your application is responsible for rotating those keys. +JWT verification requires no additional IAM permissions. Discovery and remote JWKS require outbound HTTPS; private subnets may need NAT or private connectivity. +Static `jwks` avoids network access, but your application must rotate those keys. ### Protect an HTTP route -Create `JWTVerifier` outside the Lambda handler so warm invocations reuse its signing-key cache. The verifier itself is not middleware. Calling `verifier.require()` creates middleware bound to that verifier and to the requested scopes. +`verifier.require()` creates Event Handler middleware bound to the verifier and requested scopes. Create the verifier outside the handler to reuse its cache. ```python title="middleware.py" --8<-- "examples/auth_alpha/jwt/src/middleware.py" @@ -91,13 +108,13 @@ Configure public routes and CORS preflight separately. ### Verify an Authorization header -When handling HTTP authentication without `require()`, pass the complete `Authorization` header to `verify_authorization_header()`. It validates the Bearer scheme and then calls `verify()` with the extracted JWT. +Without middleware, pass the complete `Authorization` header to `verify_authorization_header()`. It validates the Bearer scheme and verifies the extracted JWT: ```python title="direct.py" --8<-- "examples/auth_alpha/jwt/src/direct.py" ``` -Use `verify()` for an encoded JWT and `verify_authorization_header()` for the complete HTTP header. Do not split the header in application code. Both methods require `iss`, `aud`, and `exp`; `required_claims` adds more required claims. +Both verification methods require `iss`, `aud`, and `exp`; use `required_claims` to require additional claims. ## Advanced @@ -109,9 +126,10 @@ This complete Lambda adds provider-specific token checks, a required scope, a te --8<-- "examples/auth_alpha/jwt/src/custom_authorization.py" ``` -`expected_claims` must match the access-token profile documented by your identity provider. `authorize` runs only after token verification and scope checks succeed. `on_error` can change the error response and emit logs or metrics, but it never invokes the protected route. +Match `expected_claims` to your provider's access-token profile. `authorize` runs after verification and scope checks. +`on_error` can change the error response or emit logs/metrics, but cannot invoke the protected route. -The callback receives stable `reason` and `retryable` fields without token data. Preserve `error.status_code` and `error.headers` unless you intentionally want to change the HTTP contract. +The callback exposes `reason` and `retryable`, without token data. Preserve `error.status_code` and `error.headers` to retain the HTTP contract. ### Key freshness and Lambda timeouts @@ -121,7 +139,8 @@ The callback receives stable `reason` and `retryable` fields without token data. | `jwks_max_age_seconds` | 5 minutes | Limits how long fetched keys remain trusted | | `unknown_kid_cooldown_seconds` | 5 minutes | Limits repeated refreshes for unknown key IDs | -The first verification fetches signing keys unless you provide static `jwks`. Warm invocations reuse the cache. A successful refresh replaces the key set so removed keys are no longer trusted. If refresh fails after the cache expires, verification raises `JWKSFetchError` instead of using stale keys. +The first verification fetches keys unless `jwks` is static. Warm invocations reuse them; a refresh replaces the key set, dropping removed keys. +An expired cache with a failed refresh raises `JWKSFetchError`; stale keys are not used. !!! warning "Leave time for Lambda to handle the error" Set `timeout_seconds` lower than the Lambda function timeout. If both use the three-second default, Lambda can terminate the invocation before your code receives `JWKSFetchError`. @@ -132,7 +151,7 @@ Call `prefetch()` after constructing the verifier to retrieve keys during Lambda --8<-- "examples/auth_alpha/jwt/src/prefetch.py" ``` -This can reduce first-invocation latency, but an identity-provider outage can then fail the cold start. `prefetch()` is optional; without it, the first `verify()` retrieves the keys. +Prefetching can reduce first-invocation latency, but a provider outage can fail the cold start. Without it, the first `verify()` fetches keys. Static `jwks` avoids network access. Recreate the verifier or execution environment when the configured keys change. @@ -168,7 +187,8 @@ The example template disables API Gateway authorizer-result caching so every req --8<-- "examples/auth_alpha/jwt/templates/sam.yaml" ``` -If you enable Gateway caching, include all request attributes used by authorization in its identity sources to prevent decisions from being reused across different authorization inputs. Even with a complete cache key, a cached allow can outlive the JWT expiration until the cache TTL expires. Keep result caching disabled when every request must respect token expiration. This cache is independent of the verifier JWKS cache. +If enabling Gateway caching, include every authorization input in its identity sources. A cached allow can still outlive JWT expiration until the cache TTL ends. +Keep it disabled when every request must respect token expiration. Gateway caching is independent of the JWKS cache. ### Errors and diagnostics @@ -185,7 +205,7 @@ If you enable Gateway caching, include all request attributes used by authorizat | `forbidden` | false | | `jwks_unavailable` | true | -Use `reason.value` for log fields and metric dimensions. Do not parse exception messages or log tokens, claims, or request headers. `retryable=true` means a later attempt might succeed after the identity provider recovers; it does not guarantee that retrying will succeed. +Log `reason.value` and `retryable`, not exception messages, tokens, claims, or request headers. `retryable=true` means a later attempt might succeed after provider recovery. ## Testing your code @@ -195,4 +215,5 @@ Use `mock_claims` to test the complete middleware Lambda without cryptography or --8<-- "examples/auth_alpha/jwt/tests/test_middleware.py" ``` -The test still sends an HTTP API event, extracts the Bearer token, and checks the required scope before invoking the route. `mock_claims` replaces only token verification and restores the verifier when the context manager exits. Keep separate verification tests for the token profiles your application accepts. +The test exercises the HTTP event, Bearer extraction, and scope check. `mock_claims` replaces only token verification and restores it on exit. +Keep separate tests for the token profiles your application accepts. diff --git a/docs/utilities/oauth2.md b/docs/utilities/oauth2.md new file mode 100644 index 00000000000..db47dde90ce --- /dev/null +++ b/docs/utilities/oauth2.md @@ -0,0 +1,285 @@ +--- +title: OAuth2 client (alpha) +description: Acquire and cache client-credentials tokens for downstream APIs +status: new +--- + +!!! warning "Alpha / experimental" + This utility ships under the `auth_alpha` namespace while we collect feedback. Its public API may change before GA. Pin your Powertools version before using it in production. + +`OAuth2Client` obtains bearer tokens for a Lambda function calling an API on its own behalf. +It uses the client-credentials grant with `client_secret_basic` (default) or explicit `client_secret_post` authentication. + +Use [JWT verification](auth.md) to authenticate incoming requests. The OAuth client obtains separate credentials for outgoing requests; it does not forward an incoming caller's token. + +```mermaid +flowchart LR + Lambda["Lambda handler"] --> Client["OAuth2Client"] + Client -->|"No reusable token: client credentials"| IdP["Token endpoint"] + IdP -->|"Access token"| Client + Client -->|"Bearer access token"| API["Downstream API"] +``` + +## Key features + +* Cache access tokens across warm Lambda invocations and reacquire them before expiration. +* Coordinate concurrent token requests within one client. +* Resolve a client secret for each exchange attempt. +* Authenticate using HTTP Basic or form-body credentials, as required by your provider. +* Select a downstream API using a provider-specific audience or an RFC 8707 resource indicator. +* Obtain headers for your HTTP client or send a synchronous authenticated request. +* Report fixed failure reasons without exposing credentials or provider responses. + +## Getting started + +### Install + +```shell +pip install "aws-lambda-powertools[oauth2]" +``` + +The `oauth2` extra installs urllib3; PyJWT and cryptography are not required. + +#### Compatibility with boto3 + +If the function also uses boto3, including through Parameters, install and validate the SDK together with OAuth: + +```shell +pip install "aws-lambda-powertools[oauth2,aws-sdk]" +python -m pip check +``` + +Resolve all function dependencies together, lock their versions, and package boto3, botocore, and urllib3 together. Make a failed `pip check` fail the build. + +Packaged urllib3 overrides the runtime copy and can conflict with the runtime's SDK. +Build checks cannot validate dependencies supplied only by the runtime or separate layers. See [AWS packaging guidance](https://docs.aws.amazon.com/lambda/latest/dg/python-package.html#python-package-dependencies). + +### Configure your identity provider + +Register a client-credentials application with access to your API. Obtain its client ID, secret, HTTPS token endpoint, and required scopes or API identifier. +`OAuth2Client` does not register applications or discover endpoints. + +### Call a downstream API + +Create the client outside the Lambda handler so warm invocations reuse its token cache. +This complete Lambda loads a plain-string secret through Parameters and calls `GET /stock/{sku}` on an inventory API that returns JSON: + +```python title="client_credentials.py" +--8<-- "examples/auth_alpha/oauth2/src/client_credentials.py" +``` + +Configure these environment variables: + +| Variable | Required | Value | +| -------- | -------- | ----- | +| `TOKEN_URL` | Yes | Trusted HTTPS token endpoint | +| `CLIENT_ID` | Yes | Registered application client ID | +| `CLIENT_SECRET_NAME` | Yes | Secrets Manager secret name or ARN; its value must be a plain string | +| `INVENTORY_URL` | Yes | Trusted HTTPS base URL of the downstream API | +| `SCOPES` | Provider-dependent | Space-separated scopes, such as `inventory:read inventory:write`; omit when none are needed | +| `AUDIENCE` | Provider-dependent | API identifier required by your provider; omit when unused | +| `RESOURCE` | Provider-dependent | RFC 8707 resource URI; omit when unused | + +Set at most one of `AUDIENCE` and `RESOURCE`; neither is derived from `INVENTORY_URL`. See [provider configuration](#choose-the-resource). + +Invoke with `{"sku": "item/123"}`; the handler URL-encodes the SKU. + +The example allows five seconds for acquisition plus five for the API request. Set the Lambda timeout higher; see [timeouts and cold starts](#timeouts-and-cold-starts). + +### Required resources + +Parameters needs Secrets Manager connectivity, `secretsmanager:GetSecretValue`, and `kms:Decrypt` for a customer-managed key. +Local execution also needs SDK credentials and a region. OAuth token exchange itself requires no additional IAM permissions. + +Allow outbound HTTPS to the token endpoint and API. Private subnets may require NAT or private connectivity. + +### Choose the client authentication method + +| `auth_method` | Credentials sent to the token endpoint | +| ------------- | -------------------------------------- | +| `client_secret_basic` (default) | Form-encoded client ID and secret in HTTP Basic authentication; neither is in the form body | +| `client_secret_post` | `client_id` and `client_secret` in the form body, without an Authorization header | + +Select your provider's method; there is no automatic fallback. API requests always use the acquired bearer token. + +This example uses `CLIENT_SECRET` instead of `CLIENT_SECRET_NAME` and calls `INVENTORY_URL` directly: + +```python title="client_secret_post.py" +--8<-- "examples/auth_alpha/oauth2/src/client_secret_post.py" +``` + +You can also pass the Parameters `load_secret` function from the first example as `client_secret`. + +### Choose how to send the API request + +`OAuth2Client` provides two operations: + +| Method | Use when | +| ------ | -------- | +| `auth_headers()` | You want an Authorization header for an application-owned HTTP client | +| `request(method, url, ...)` | You want the utility to send an authenticated HTTPS request | + +Both acquire tokens on demand; construction fetches no secret or token. + +### Choose the resource + +| Parameter | Token-request field | Purpose | +| --------- | ------------------- | ------- | +| `audience` | `audience=` | Provider-specific API selection, such as an Auth0 API identifier | +| `resource` | `resource=` | One resource indicator for providers supporting RFC 8707 | + +Use at most one parameter, as required by your provider. If omitted, API selection depends on the provider's client configuration or scope conventions. + +`resource` must be an absolute URI without a fragment, such as `https://inventory.example.com` or `urn:example:inventory`. +`audience` accepts a provider-specific, nonempty string. Both are validated during construction. + +Common provider configurations: + +| Provider | Token endpoint path | API selection | +| -------- | ------------------- | ------------- | +| [Amazon Cognito](https://docs.aws.amazon.com/cognito/latest/developerguide/token-endpoint.html) | `/oauth2/token` on the user pool domain | Custom resource-server scopes such as `inventory/read` | +| [Auth0](https://auth0.com/docs/get-started/authentication-and-authorization-flow/client-credentials-flow/call-your-api-using-the-client-credentials-flow) | `/oauth/token` | `audience` set to the API identifier; scopes as configured for the API | +| [Microsoft Entra ID](https://learn.microsoft.com/en-us/entra/identity-platform/v2-oauth2-client-creds-grant-flow) | `/{tenant}/oauth2/v2.0/token` | One resource's `/.default` scope, such as `https://graph.microsoft.com/.default` | +| [Okta custom authorization server](https://developer.okta.com/docs/guides/implement-grant-type/clientcreds/main/) | `/oauth2/{authorizationServerId}/v1/token` | Custom API scopes | +| [Keycloak](https://www.keycloak.org/docs/latest/server_admin/index.html#_service_accounts) | `/realms/{realm}/protocol/openid-connect/token` | Service-account roles and client scopes configured in the realm | + +Okta's organization authorization server [requires `private_key_jwt` for service apps](https://developer.okta.com/docs/guides/implement-oauth-for-okta-serviceapp/main/) requesting Okta management scopes. That authentication method is outside this client's scope. + +Create one client per provider and resource. The request URL does not change token selection; changing the resource requires a new client. +Your application selects the client; routing and failover are not automatic. + +### Use your own HTTP client + +Pass the fresh `Authorization: Bearer ` dictionary from `auth_headers()` to your HTTP client: + +```python title="headers.py" +--8<-- "examples/auth_alpha/oauth2/src/headers.py" +``` + +This example validates HTTPS before acquiring credentials. Configure your HTTP client's timeouts, redirects, and retries yourself. +Never log the returned headers or send them to an untrusted destination. + +## Advanced + +### Lambda execution environments and token lifetimes + +Each client caches tokens within one Lambda execution environment. Warm invocations may reuse them; new environments acquire their own tokens and can exchange concurrently. + +Cached tokens are reused while more than 30 seconds remain. After that, the next call performs a new client-credentials exchange, not a refresh-token grant. +Tokens with missing or at most 30 seconds of advertised lifetime are returned uncached. Invalid lifetimes and tokens that expire during acquisition are rejected. + +Lifetime accounting starts immediately before the token request. Secret lookup consumes the acquisition budget without shortening the token's lifetime. + +Concurrent callers on one client share an exchange, including short-lived tokens and failures. Each waiting caller keeps its own acquisition deadline. + +### Secret rotation + +`client_secret` accepts a nonempty string or a callable, invoked for each exchange attempt including retries. The client does not cache the callable's result. + +Cached access tokens may remain usable after rotation. Parameters' `max_age=300` can delay reading an updated secret by five minutes. + +Environment-variable secrets are static; a callable reading `os.environ` does not fetch external updates. + +### Timeouts and cold starts + +| Setting | Default | Applies to | +| ------- | ------- | ---------- | +| `OAuth2Client(timeout_seconds=...)` | 3 seconds | One acquisition, including waiting, secret lookup, token requests, and retry backoff | +| `request(..., timeout=...)` | 5 seconds | The downstream request after token acquisition, including reading its response | +| Secret-provider SDK timeouts | Provider-specific | Each secret lookup; configure independently | + +These budgets are sequential. Set the Lambda timeout above their sum, leaving time for application work and error handling. + +Measure cold starts at your chosen memory and network settings: initial secret lookup and SDK setup can exceed the three-second default. +The Parameters example initializes the SDK outside the handler, uses one-second connect/two-second read timeouts with one SDK attempt, and allows five seconds for acquisition. Tune these values for your workload. + +HTTP header and body reads use the remaining budget. DNS resolution, secret loaders, and upload producers cannot always be interrupted, although their elapsed time still counts. +Configure their own timeouts where supported. + +### Retries and downstream responses + +Token-endpoint transport failures, HTTP 429, and HTTP 5xx allow at most two retries, with 100 ms then 200 ms backoff within the acquisition budget. +Other HTTP errors, malformed token responses, and secret-loader failures are not retried. + +`request()` returns an urllib3 HTTP response with `.status`, `.headers`, `.data`, and `.json()`. +Handle HTTP statuses in your application: even 401/403 does not invalidate the cached token or trigger another exchange. Downstream operations are never automatically replayed. + +Responses are fully buffered in memory; use `auth_headers()` and a streaming client for large downloads. +The examples expect HTTP 200 with JSON. Handle other statuses or empty bodies according to your API. + +### Destination safety + +`request()` requires HTTPS, disables redirects, and forwards only `body`, `fields`, `json`, `encode_multipart`, and `multipart_boundary` options to urllib3. +Use `auth_headers()` for other transport options. + +Header names must follow HTTP token syntax; values must fit Latin-1 without ASCII controls other than tabs. +Invalid headers and case-insensitive Authorization overrides are rejected before secret lookup. + +!!! warning "Use trusted destination URLs" + Use configured, trusted URLs, never caller-controlled destinations. `request()` does not restrict URLs using the configured audience or resource. + +### Errors and diagnostics + +OAuth errors inherit from `AuthError` in `auth_alpha.exceptions`. + +| Exception | Reason | Retryable | +| --------- | ------ | --------- | +| `TokenExchangeError` | `token_exchange_failed` | True for transient endpoint failures or acquisition timeouts; otherwise false | +| `DownstreamRequestError` | `downstream_request_failed` | False: the server may already have performed the operation | + +Invalid client configuration, method, URL, timeout, headers, or option names raise `ValueError` before token acquisition. +`retryable=true` means a later acquisition might succeed. + +Use the fixed `reason.value` and `retryable` fields for logs and metrics: + +```python title="diagnostics.py" +--8<-- "examples/auth_alpha/oauth2/src/diagnostics.py" +``` + +This handler returns proxy-style responses and uses the environment-secret deployment settings. + +The client emits no logs and removes provider exception chains. Never log secrets, tokens, request headers/bodies, or provider responses. + +### Supported scope + +The client treats bearer tokens, including JWTs, as opaque strings. The downstream API validates and authorizes them. + +Other grants, introspection, revocation, JWT client authentication, mTLS, and DPoP are not supported. Calls are synchronous. + +### Calling downstream APIs from an MCP tool + +After authorizing the MCP caller, use `asyncio.to_thread()` to call the downstream API with separate client credentials: + +```python +import asyncio +from urllib.parse import quote + +from client_credentials import INVENTORY_URL, inventory_api +from mcp.server.auth.middleware.auth_context import get_access_token + + +async def check_stock(sku: str) -> dict: + caller = get_access_token() + if caller is None or "inventory:read" not in caller.scopes: + raise PermissionError("Inventory read permission is required") + response = await asyncio.to_thread( + inventory_api.request, + "GET", + f"{INVENTORY_URL}/stock/{quote(sku, safe='')}", + timeout=5, + ) + if response.status != 200: + raise RuntimeError("Inventory lookup failed") + return response.json() +``` + +Package `client_credentials.py` alongside the tool and register `check_stock` with an authenticated MCP server. The SDK owns transport authentication and error responses. +Cancelling the task does not stop the worker's request; timeouts still apply. Install the MCP SDK separately. + +## Testing your code + +Mock the client operation when testing application behavior, and test your provider configuration separately: + +```python title="test_client_credentials.py" +--8<-- "examples/auth_alpha/oauth2/tests/test_client_credentials.py" +``` diff --git a/examples/auth_alpha/oauth2/src/client_credentials.py b/examples/auth_alpha/oauth2/src/client_credentials.py new file mode 100644 index 00000000000..036f2192133 --- /dev/null +++ b/examples/auth_alpha/oauth2/src/client_credentials.py @@ -0,0 +1,44 @@ +import os +from urllib.parse import quote + +from botocore.config import Config + +from aws_lambda_powertools.utilities import parameters +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.typing import LambdaContext + +secrets = parameters.SecretsProvider( + boto_config=Config( + connect_timeout=1, + read_timeout=2, + retries={"total_max_attempts": 1}, + ), +) + + +def load_secret() -> str: + secret = secrets.get(os.environ["CLIENT_SECRET_NAME"], max_age=300) + if not isinstance(secret, str): + raise ValueError("Expected a string client secret") + return secret + + +INVENTORY_URL = os.environ["INVENTORY_URL"] + +inventory_api = OAuth2Client( + token_url=os.environ["TOKEN_URL"], + client_id=os.environ["CLIENT_ID"], + client_secret=load_secret, + scopes=os.environ.get("SCOPES", "").split(), + audience=os.environ.get("AUDIENCE"), + resource=os.environ.get("RESOURCE"), + timeout_seconds=5, +) + + +def lambda_handler(event: dict, context: LambdaContext): + sku = quote(event["sku"], safe="") + response = inventory_api.request("GET", f"{INVENTORY_URL}/stock/{sku}", timeout=5) + if response.status != 200: + raise RuntimeError("Inventory lookup failed") + return response.json() diff --git a/examples/auth_alpha/oauth2/src/client_secret_post.py b/examples/auth_alpha/oauth2/src/client_secret_post.py new file mode 100644 index 00000000000..9d2dd112c48 --- /dev/null +++ b/examples/auth_alpha/oauth2/src/client_secret_post.py @@ -0,0 +1,25 @@ +import os + +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.typing import LambdaContext + +# Create outside the handler so warm invocations can reuse the token cache. +# A callable client_secret loader is also supported, as in client_credentials.py. +inventory_api = OAuth2Client( + token_url=os.environ["TOKEN_URL"], + client_id=os.environ["CLIENT_ID"], + client_secret=os.environ["CLIENT_SECRET"], + auth_method="client_secret_post", + scopes=os.environ.get("SCOPES", "").split(), + audience=os.environ.get("AUDIENCE"), + resource=os.environ.get("RESOURCE"), + timeout_seconds=5, +) +INVENTORY_URL = os.environ["INVENTORY_URL"] + + +def lambda_handler(event: dict, context: LambdaContext): + response = inventory_api.request("GET", INVENTORY_URL, timeout=5) + if response.status != 200: + raise RuntimeError("Inventory lookup failed") + return response.json() diff --git a/examples/auth_alpha/oauth2/src/diagnostics.py b/examples/auth_alpha/oauth2/src/diagnostics.py new file mode 100644 index 00000000000..d1c41daffc8 --- /dev/null +++ b/examples/auth_alpha/oauth2/src/diagnostics.py @@ -0,0 +1,32 @@ +import json +import os +from urllib.parse import quote + +from aws_lambda_powertools import Logger +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions import DownstreamRequestError, TokenExchangeError +from aws_lambda_powertools.utilities.typing import LambdaContext + +logger = Logger() +INVENTORY_URL = os.environ["INVENTORY_URL"] +inventory_api = OAuth2Client( + token_url=os.environ["TOKEN_URL"], + client_id=os.environ["CLIENT_ID"], + client_secret=lambda: os.environ["CLIENT_SECRET"], + scopes=os.environ.get("SCOPES", "").split(), + audience=os.environ.get("AUDIENCE"), + resource=os.environ.get("RESOURCE"), + timeout_seconds=5, +) + + +def lambda_handler(event: dict, context: LambdaContext): + sku = quote(event["sku"], safe="") + try: + response = inventory_api.request("GET", f"{INVENTORY_URL}/stock/{sku}", timeout=5) + except (TokenExchangeError, DownstreamRequestError) as error: + logger.warning("Inventory request unavailable", reason=error.reason.value, retryable=error.retryable) + return {"statusCode": 502, "body": "Inventory request unavailable"} + if response.status != 200: + return {"statusCode": 502, "body": "Inventory request unavailable"} + return {"statusCode": 200, "body": json.dumps(response.json())} diff --git a/examples/auth_alpha/oauth2/src/headers.py b/examples/auth_alpha/oauth2/src/headers.py new file mode 100644 index 00000000000..42d4022d3fd --- /dev/null +++ b/examples/auth_alpha/oauth2/src/headers.py @@ -0,0 +1,44 @@ +import os +from urllib.parse import quote, urlsplit + +import urllib3 + +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.typing import LambdaContext + +INVENTORY_URL = os.environ["INVENTORY_URL"] +inventory_url = urlsplit(INVENTORY_URL) +if ( + inventory_url.scheme != "https" + or not inventory_url.hostname + or inventory_url.username is not None + or inventory_url.password is not None + or "#" in INVENTORY_URL +): + raise ValueError("INVENTORY_URL must be an HTTPS URL without user information or a fragment") + +inventory_api = OAuth2Client( + token_url=os.environ["TOKEN_URL"], + client_id=os.environ["CLIENT_ID"], + client_secret=lambda: os.environ["CLIENT_SECRET"], + scopes=os.environ.get("SCOPES", "").split(), + audience=os.environ.get("AUDIENCE"), + resource=os.environ.get("RESOURCE"), + timeout_seconds=5, +) +http = urllib3.PoolManager() + + +def lambda_handler(event: dict, context: LambdaContext): + sku = quote(event["sku"], safe="") + response = http.request( + "GET", + f"{INVENTORY_URL}/stock/{sku}", + headers=inventory_api.auth_headers(), + timeout=urllib3.Timeout(total=5), + redirect=False, + retries=False, + ) + if response.status != 200: + raise RuntimeError("Inventory lookup failed") + return response.json() diff --git a/examples/auth_alpha/oauth2/tests/test_client_credentials.py b/examples/auth_alpha/oauth2/tests/test_client_credentials.py new file mode 100644 index 00000000000..0558234761d --- /dev/null +++ b/examples/auth_alpha/oauth2/tests/test_client_credentials.py @@ -0,0 +1,20 @@ +import urllib3 + + +def test_inventory_lookup(monkeypatch, mocker): + monkeypatch.setenv("TOKEN_URL", "https://idp.example.com/token") + monkeypatch.setenv("CLIENT_ID", "orders") + monkeypatch.setenv("CLIENT_SECRET_NAME", "orders/oauth-secret") + monkeypatch.setenv("INVENTORY_URL", "https://inventory.example.com") + mocker.patch("aws_lambda_powertools.utilities.parameters.SecretsProvider") + + from client_credentials import inventory_api, lambda_handler + + def request(method, url, *, timeout): + assert method == "GET" + assert url == "https://inventory.example.com/stock/item%2F123" + assert timeout == 5 + return urllib3.HTTPResponse(body=b'{"stock":12}', status=200) + + monkeypatch.setattr(inventory_api, "request", request) + assert lambda_handler({"sku": "item/123"}, {}) == {"stock": 12} diff --git a/examples/auth_alpha/oauth2/tests/test_client_secret_post.py b/examples/auth_alpha/oauth2/tests/test_client_secret_post.py new file mode 100644 index 00000000000..3f448700240 --- /dev/null +++ b/examples/auth_alpha/oauth2/tests/test_client_secret_post.py @@ -0,0 +1,47 @@ +import importlib +import sys +from io import BytesIO +from urllib.parse import parse_qs + +import urllib3 + + +def test_post_credentials_are_sent_only_to_the_token_endpoint(monkeypatch): + monkeypatch.setenv("TOKEN_URL", "https://idp.example.com/token") + monkeypatch.setenv("CLIENT_ID", "orders") + monkeypatch.setenv("CLIENT_SECRET", "test-only-secret") + monkeypatch.setenv("INVENTORY_URL", "https://inventory.example.com") + monkeypatch.setenv("SCOPES", "inventory:read") + monkeypatch.delenv("AUDIENCE", raising=False) + monkeypatch.delenv("RESOURCE", raising=False) + calls = [] + + def request(self, method, url, **options): + calls.append(url) + if url == "https://idp.example.com/token": + assert method == "POST" + assert options["headers"] == {"Content-Type": "application/x-www-form-urlencoded"} + assert parse_qs(options["body"].decode()) == { + "grant_type": ["client_credentials"], + "scope": ["inventory:read"], + "client_id": ["orders"], + "client_secret": ["test-only-secret"], + } + return urllib3.HTTPResponse( + body=BytesIO(b'{"access_token":"test-token","token_type":"Bearer","expires_in":600}'), + status=200, + preload_content=False, + ) + assert method == "GET" + assert url == "https://inventory.example.com" + assert options["headers"] == {"Authorization": "Bearer test-token"} + assert "test-only-secret" not in str(options) + return urllib3.HTTPResponse(body=BytesIO(b'{"stock":12}'), status=200, preload_content=False) + + monkeypatch.setattr(urllib3.PoolManager, "request", request) + monkeypatch.delitem(sys.modules, "client_secret_post", raising=False) + example = importlib.import_module("client_secret_post") + + assert example.lambda_handler({}, {}) == {"stock": 12} + assert example.lambda_handler({}, {}) == {"stock": 12} + assert calls == ["https://idp.example.com/token", "https://inventory.example.com", "https://inventory.example.com"] diff --git a/examples/auth_alpha/oauth2/tests/test_documented_configuration.py b/examples/auth_alpha/oauth2/tests/test_documented_configuration.py new file mode 100644 index 00000000000..85867570dbc --- /dev/null +++ b/examples/auth_alpha/oauth2/tests/test_documented_configuration.py @@ -0,0 +1,118 @@ +import base64 +import json +import runpy +from io import BytesIO +from pathlib import Path +from urllib.parse import parse_qs + +import pytest +import urllib3 + +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions import DownstreamRequestError, TokenExchangeError + +EXAMPLES = Path(__file__).parents[1] / "src" + + +@pytest.fixture +def deployment(monkeypatch, mocker): + for name, value in { + "TOKEN_URL": "https://idp.example.com/token", + "CLIENT_ID": "orders", + "CLIENT_SECRET": "test-secret", + "CLIENT_SECRET_NAME": "orders/oauth-secret", + "INVENTORY_URL": "https://inventory.example.com", + }.items(): + monkeypatch.setenv(name, value) + for name in ("SCOPES", "AUDIENCE", "RESOURCE"): + monkeypatch.delenv(name, raising=False) + provider = mocker.patch("aws_lambda_powertools.utilities.parameters.SecretsProvider") + provider.return_value.get.return_value = "test-secret" + return provider + + +@pytest.mark.parametrize("filename", ["client_credentials.py", "client_secret_post.py", "headers.py", "diagnostics.py"]) +@pytest.mark.parametrize( + "configuration", + [ + {}, + {"SCOPES": "inventory/read"}, + {"SCOPES": "inventory:read inventory:write", "AUDIENCE": "inventory-api"}, + {"SCOPES": "https://graph.microsoft.com/.default"}, + {"SCOPES": "inventory:read", "RESOURCE": "urn:example:inventory"}, + ], + ids=["no-scopes", "cognito-scopes", "auth0-audience", "entra-default-scope", "resource-uri"], +) +def test_examples_send_the_configured_scopes_and_resource(deployment, monkeypatch, filename, configuration): + for name, value in configuration.items(): + monkeypatch.setenv(name, value) + calls = [] + expected_fields = {"grant_type": ["client_credentials"]} + for variable, field in (("SCOPES", "scope"), ("AUDIENCE", "audience"), ("RESOURCE", "resource")): + if variable in configuration: + expected_fields[field] = [configuration[variable]] + + def request(self, method, url, **options): + calls.append((method, url)) + if url == "https://idp.example.com/token": + assert method == "POST" + assert options["headers"]["Content-Type"] == "application/x-www-form-urlencoded" + fields = parse_qs(options["body"].decode()) + if filename == "client_secret_post.py": + assert "Authorization" not in options["headers"] + assert fields.pop("client_id") == ["orders"] + assert fields.pop("client_secret") == ["test-secret"] + else: + basic = options["headers"]["Authorization"].removeprefix("Basic ") + assert base64.b64decode(basic).decode() == "orders:test-secret" + assert fields == expected_fields + return urllib3.HTTPResponse( + body=BytesIO(b'{"access_token":"test-token","token_type":"Bearer","expires_in":600}'), + status=200, + preload_content=False, + ) + assert method == "GET" + suffix = "" if filename == "client_secret_post.py" else "/stock/item%2F123" + assert url == f"https://inventory.example.com{suffix}" + assert options["headers"] == {"Authorization": "Bearer test-token"} + assert "test-secret" not in str(options) + return urllib3.HTTPResponse(body=BytesIO(b'{"stock":12}'), status=200, preload_content=False) + + monkeypatch.setattr(urllib3.PoolManager, "request", request) + example = runpy.run_path(str(EXAMPLES / filename)) + assert calls == [] + + for _ in range(2): + response = example["lambda_handler"]({"sku": "item/123"}, {}) + if filename == "diagnostics.py": + assert response["statusCode"] == 200 + assert json.loads(response["body"]) == {"stock": 12} + else: + assert response == {"stock": 12} + assert len(calls) == 3 + if filename == "client_credentials.py": + provider = deployment + provider.return_value.get.assert_called_once_with("orders/oauth-secret", max_age=300) + config = provider.call_args.kwargs["boto_config"] + assert config.connect_timeout == 1 + assert config.read_timeout == 2 + assert config.retries == {"total_max_attempts": 1} + + +@pytest.mark.parametrize( + "failure", + [TokenExchangeError(retryable=True), DownstreamRequestError(), 403], + ids=["token-unavailable", "downstream-unavailable", "api-forbidden"], +) +def test_diagnostics_returns_an_http_error_without_replaying_the_request(deployment, monkeypatch, failure, mocker): + if isinstance(failure, Exception): + request = mocker.Mock(side_effect=failure) + else: + request = mocker.Mock(return_value=urllib3.HTTPResponse(status=failure)) + monkeypatch.setattr(OAuth2Client, "request", request) + example = runpy.run_path(str(EXAMPLES / "diagnostics.py")) + + response = example["lambda_handler"]({"sku": "item/123"}, {}) + + assert response == {"statusCode": 502, "body": "Inventory request unavailable"} + request.assert_called_once_with("GET", "https://inventory.example.com/stock/item%2F123", timeout=5) diff --git a/examples/auth_alpha/oauth2/tests/test_headers.py b/examples/auth_alpha/oauth2/tests/test_headers.py new file mode 100644 index 00000000000..64811c4df8e --- /dev/null +++ b/examples/auth_alpha/oauth2/tests/test_headers.py @@ -0,0 +1,70 @@ +import runpy +from pathlib import Path + +import pytest +import urllib3 + +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client + +EXAMPLE = Path(__file__).parents[1] / "src" / "headers.py" + + +@pytest.fixture +def example_environment(monkeypatch): + monkeypatch.setenv("TOKEN_URL", "https://idp.example.com/token") + monkeypatch.setenv("CLIENT_ID", "orders") + monkeypatch.setenv("CLIENT_SECRET", "test-secret") + + +@pytest.mark.parametrize( + "destination", + [ + "http://inventory.example.com", + "//inventory.example.com", + "https:///inventory", + "https://user:password@inventory.example.com", + "https://inventory.example.com/#fragment", + ], +) +def test_unsafe_destinations_are_rejected_before_obtaining_headers(example_environment, monkeypatch, destination): + monkeypatch.setenv("INVENTORY_URL", destination) + acquired = [] + requests = [] + + def auth_headers(self): + acquired.append(True) + return {"Authorization": "Bearer test-token"} + + def request(self, method, url, **options): + requests.append(url) + return urllib3.HTTPResponse(body=b'{"stock":12}', status=200) + + monkeypatch.setattr(OAuth2Client, "auth_headers", auth_headers) + monkeypatch.setattr(urllib3.PoolManager, "request", request) + example_path = str(EXAMPLE) + + with pytest.raises(ValueError, match="INVENTORY_URL"): + runpy.run_path(example_path) + + assert acquired == [] + assert requests == [] + + +def test_https_destination_receives_the_token_with_safe_transport_options(example_environment, monkeypatch): + monkeypatch.setenv("INVENTORY_URL", "https://inventory.example.com") + monkeypatch.setattr(OAuth2Client, "auth_headers", lambda self: {"Authorization": "Bearer test-token"}) + requests = [] + + def request(self, method, url, **options): + requests.append((method, url)) + assert options["headers"] == {"Authorization": "Bearer test-token"} + assert options["timeout"].total == 5 + assert options["redirect"] is False + assert options["retries"] is False + return urllib3.HTTPResponse(body=b'{"stock":12}', status=200) + + monkeypatch.setattr(urllib3.PoolManager, "request", request) + example = runpy.run_path(str(EXAMPLE)) + + assert example["lambda_handler"]({"sku": "item/123"}, {}) == {"stock": 12} + assert requests == [("GET", "https://inventory.example.com/stock/item%2F123")] diff --git a/mkdocs.yml b/mkdocs.yml index b9f02a668b8..7321f451d83 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -26,7 +26,9 @@ nav: - core/event_handler/appsync_events.md - core/event_handler/bedrock_agents.md - utilities/parameters.md - - Auth (alpha): utilities/auth.md + - Auth (alpha): + - JWT verification: utilities/auth.md + - OAuth2 client: utilities/oauth2.md - utilities/batch.md - utilities/kafka.md - utilities/typing.md @@ -250,6 +252,7 @@ plugins: - core/event_handler/bedrock_agents.md Utilities: - utilities/auth.md + - utilities/oauth2.md - utilities/parameters.md - utilities/batch.md - utilities/typing.md diff --git a/noxfile.py b/noxfile.py index 1514e93b444..4324fd96eb9 100644 --- a/noxfile.py +++ b/noxfile.py @@ -235,6 +235,16 @@ def test_with_auth_required_packages(session: nox.Session): """Verify the Auth utility using only its declared optional dependencies.""" build_and_run_test( session, - folders=[f"{PREFIX_TESTS_FUNCTIONAL}/auth_alpha/"], + folders=[f"{PREFIX_TESTS_FUNCTIONAL}/auth_alpha/jwt/"], extras="jwt", ) + + +@nox.session() +def test_with_oauth2_required_packages(session: nox.Session): + """Verify OAuth token acquisition without JWT or cryptography dependencies.""" + build_and_run_test( + session, + folders=[f"{PREFIX_TESTS_FUNCTIONAL}/auth_alpha/oauth2/"], + extras="oauth2", + ) diff --git a/pyproject.toml b/pyproject.toml index d19bb94671c..ab2041242e0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -64,6 +64,9 @@ jwt = [ "cryptography (>=50.0.1,<51.0.0)", "urllib3 (>=2.8.0,<3.0.0)", ] +oauth2 = [ + "urllib3 (>=2.8.0,<3.0.0)", +] all = [ "pydantic (>=2.4.0,<3.0.0)", "pydantic-settings (>=2.6.1,<3.0.0)", diff --git a/tests/functional/auth_alpha/conftest.py b/tests/functional/auth_alpha/conftest.py new file mode 100644 index 00000000000..f6873ffbc10 --- /dev/null +++ b/tests/functional/auth_alpha/conftest.py @@ -0,0 +1,57 @@ +import io +import json +import time +from collections import deque + +import pytest +import urllib3 + + +class FakeHTTP: + """In-memory authentication endpoints at the HTTP transport boundary.""" + + def __init__(self): + self.responses = {} + self.requests = [] + + def serve(self, url, body, *, status=200, method="GET"): + self.responses[(method, url)] = deque([(status, body)]) + + def request(self, method, url, **kwargs): + self.requests.append((method, url, kwargs)) + responses = self.responses[(method, url)] + status, body = responses[0] if len(responses) == 1 else responses.popleft() + if callable(body): + body = body() + if isinstance(body, Exception): + raise body + payload = body if isinstance(body, bytes) else json.dumps(body).encode() + return urllib3.HTTPResponse( + body=io.BytesIO(payload), + headers={"content-type": "application/json"}, + status=status, + preload_content=False, + ) + + +@pytest.fixture +def http(monkeypatch): + transport = FakeHTTP() + monkeypatch.setattr(urllib3, "PoolManager", lambda **kwargs: transport) + return transport + + +@pytest.fixture +def clock(monkeypatch): + class Clock: + now = 1000.0 + + def __call__(self): + return self.now + + def advance(self, seconds): + self.now += seconds + + clock = Clock() + monkeypatch.setattr(time, "monotonic", clock) + return clock diff --git a/tests/functional/auth_alpha/jwt/conftest.py b/tests/functional/auth_alpha/jwt/conftest.py index 69c67d98155..e5259dab923 100644 --- a/tests/functional/auth_alpha/jwt/conftest.py +++ b/tests/functional/auth_alpha/jwt/conftest.py @@ -1,12 +1,8 @@ -import io -import json import time import weakref -from collections import deque import jwt import pytest -import urllib3 from cryptography.hazmat.primitives.asymmetric import rsa from aws_lambda_powertools.utilities.auth_alpha.jwt._internal import jwks as jwks_module @@ -47,54 +43,7 @@ def issue(payload=None, *, key=None, kid="key-1", algorithm="RS256", headers=Non return issue -class FakeHTTP: - """In-memory JWKS endpoints at the HTTP transport boundary.""" - - def __init__(self): - self.responses = {} - self.requests = [] - - def serve(self, url, body, *, status=200, method="GET"): - self.responses[(method, url)] = deque([(status, body)]) - - def request(self, method, url, **kwargs): - self.requests.append((method, url, kwargs)) - responses = self.responses[(method, url)] - status, body = responses[0] if len(responses) == 1 else responses.popleft() - if callable(body): - body = body() - if isinstance(body, Exception): - raise body - payload = body if isinstance(body, bytes) else json.dumps(body).encode() - return urllib3.HTTPResponse( - body=io.BytesIO(payload), - headers={"content-type": "application/json"}, - status=status, - preload_content=False, - ) - - -@pytest.fixture -def http(monkeypatch): - # Each fake provider belongs to one test. Error tracebacks can keep a - # previous verifier alive; retain sharing only within the current test. +@pytest.fixture(autouse=True) +def isolated_jwks_caches(monkeypatch): + # Each test's fake provider owns its cache, independent of retained tracebacks. monkeypatch.setattr(jwks_module, "_caches", weakref.WeakValueDictionary()) - transport = FakeHTTP() - monkeypatch.setattr(urllib3, "PoolManager", lambda **kwargs: transport) - return transport - - -@pytest.fixture -def clock(monkeypatch): - class Clock: - now = 1000.0 - - def __call__(self): - return self.now - - def advance(self, seconds): - self.now += seconds - - clock = Clock() - monkeypatch.setattr(time, "monotonic", clock) - return clock diff --git a/tests/functional/auth_alpha/oauth2/__init__.py b/tests/functional/auth_alpha/oauth2/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/functional/auth_alpha/oauth2/test_auth_methods.py b/tests/functional/auth_alpha/oauth2/test_auth_methods.py new file mode 100644 index 00000000000..cbbac878978 --- /dev/null +++ b/tests/functional/auth_alpha/oauth2/test_auth_methods.py @@ -0,0 +1,124 @@ +import base64 +import time +from collections import deque +from urllib.parse import parse_qs, unquote_plus + +import pytest + +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions import TokenExchangeError + +TOKEN_URL = "https://idp.example.com/token?tenant=inventory" +TOKEN_RESPONSE = {"access_token": "test-token", "token_type": "Bearer", "expires_in": 600} + + +@pytest.mark.parametrize( + "options", + [{}, {"auth_method": "client_secret_basic"}, {"auth_method": "client_secret_post"}], + ids=["default-basic", "explicit-basic", "post"], +) +@pytest.mark.parametrize("selection", [{}, {"audience": "inventory"}, {"resource": "urn:example:inventory"}]) +@pytest.mark.parametrize( + ("client_id", "secret"), + [ + ("client:id", "secret:value with space+&=%25"), + ("clïent:日本語", "sécret/日本語&scope=admin\r\ninjected"), + ], +) +def test_credentials_use_only_the_selected_location_and_are_encoded_once(http, options, selection, client_id, secret): + http.serve(TOKEN_URL, TOKEN_RESPONSE, method="POST") + subject = OAuth2Client( + token_url=TOKEN_URL, + client_id=client_id, + client_secret=secret, + scopes=["inventory:read", "inventory:write"], + **selection, + **options, + ) + + assert subject.auth_headers() == {"Authorization": "Bearer test-token"} + assert subject.auth_headers() == {"Authorization": "Bearer test-token"} + + assert len(http.requests) == 1 + method, url, request = http.requests[0] + assert method == "POST" + assert url == TOKEN_URL + assert request["headers"]["Content-Type"] == "application/x-www-form-urlencoded" + expected = { + "grant_type": ["client_credentials"], + "scope": ["inventory:read inventory:write"], + **{key: [value] for key, value in selection.items()}, + } + if options.get("auth_method") == "client_secret_post": + expected.update(client_id=[client_id], client_secret=[secret]) + assert "Authorization" not in request["headers"] + else: + scheme, authorization = request["headers"]["Authorization"].split(" ", 1) + assert scheme == "Basic" + encoded_id, encoded_secret = base64.b64decode(authorization, validate=True).decode().split(":") + assert unquote_plus(encoded_id) == client_id + assert unquote_plus(encoded_secret) == secret + assert parse_qs(request["body"].decode("ascii"), keep_blank_values=True) == expected + + +@pytest.mark.parametrize( + "auth_method", + [None, "", "basic", "post", "CLIENT_SECRET_POST", "client_secret_jwt", "private_key_jwt", "none", [], {}, 1], +) +def test_invalid_auth_method_is_rejected_before_loading_credentials(http, mocker, auth_method): + load_secret = mocker.Mock(return_value="private-test-secret") + with pytest.raises(ValueError, match="auth_method"): + OAuth2Client( + token_url=TOKEN_URL, + client_id="client", + client_secret=load_secret, + auth_method=auth_method, + ) + + load_secret.assert_not_called() + assert http.requests == [] + + +@pytest.mark.parametrize("status", [429, 503]) +def test_post_rebuilds_the_form_with_the_current_secret_on_retry(http, clock, monkeypatch, status): + secrets = deque(["initial secret+&=", "rotated secret+&="]) + http.responses[("POST", TOKEN_URL)] = deque([(status, {}), (200, TOKEN_RESPONSE)]) + monkeypatch.setattr(time, "sleep", clock.advance) + subject = OAuth2Client( + token_url=TOKEN_URL, + client_id="orders", + client_secret=secrets.popleft, + auth_method="client_secret_post", + scopes=["inventory:read"], + resource="urn:example:inventory", + ) + + assert subject.auth_headers() == {"Authorization": "Bearer test-token"} + assert not secrets + assert len(http.requests) == 2 + for request, secret in zip(http.requests, ["initial secret+&=", "rotated secret+&="], strict=True): + assert request[1] == TOKEN_URL + assert request[2]["headers"] == {"Content-Type": "application/x-www-form-urlencoded"} + assert parse_qs(request[2]["body"].decode()) == { + "grant_type": ["client_credentials"], + "scope": ["inventory:read"], + "resource": ["urn:example:inventory"], + "client_id": ["orders"], + "client_secret": [secret], + } + + +@pytest.mark.parametrize("auth_method", ["client_secret_basic", "client_secret_post"]) +def test_authentication_rejection_does_not_fall_back_to_another_method(http, auth_method): + http.serve(TOKEN_URL, {"error": "invalid_client"}, method="POST", status=401) + subject = OAuth2Client( + token_url=TOKEN_URL, + client_id="client", + client_secret="test-secret", + auth_method=auth_method, + ) + + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + assert not error.value.retryable + assert len(http.requests) == 1 diff --git a/tests/functional/auth_alpha/oauth2/test_client.py b/tests/functional/auth_alpha/oauth2/test_client.py new file mode 100644 index 00000000000..5b988a91c57 --- /dev/null +++ b/tests/functional/auth_alpha/oauth2/test_client.py @@ -0,0 +1,741 @@ +import base64 +import threading +import time +import traceback +from collections import deque +from concurrent.futures import ThreadPoolExecutor +from urllib.parse import parse_qs + +import pytest + +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions import TokenExchangeError + +TOKEN_URL = "https://idp.example.com/oauth/token" + + +@pytest.fixture(params=["client_secret_basic", "client_secret_post"]) +def auth_method(request): + return request.param + + +def client(**options): + config = { + "token_url": TOKEN_URL, + "client_id": "orders-client", + "client_secret": "test-client-secret", + "scopes": ["orders:read"], + } + return OAuth2Client(**{**config, **options}) + + +def test_client_credentials_exchange_selects_resource_and_caches_token(http): + http.serve( + TOKEN_URL, + {"access_token": "opaque-access-token", "token_type": "Bearer", "expires_in": 3600}, + method="POST", + ) + subject = client(audience="https://api.example.com") + + assert subject.auth_headers() == {"Authorization": "Bearer opaque-access-token"} + assert subject.auth_headers() == {"Authorization": "Bearer opaque-access-token"} + assert len(http.requests) == 1 + method, url, request = http.requests[0] + assert method == "POST" + assert url == TOKEN_URL + assert parse_qs(request["body"].decode()) == { + "grant_type": ["client_credentials"], + "scope": ["orders:read"], + "audience": ["https://api.example.com"], + } + assert request["headers"]["Content-Type"] == "application/x-www-form-urlencoded" + + +def test_basic_auth_encodes_each_credential_before_base64(http): + http.serve(TOKEN_URL, {"access_token": "token", "token_type": "bearer", "expires_in": 3600}, method="POST") + subject = client(client_id="client:id", client_secret="secret:value with space") + subject.auth_headers() + request = http.requests[0][2] + encoded = request["headers"]["Authorization"].removeprefix("Basic ") + + assert base64.b64decode(encoded).decode() == "client%3Aid:secret%3Avalue+with+space" + assert "client_secret" not in parse_qs(request["body"].decode()) + + +def test_token_is_reacquired_before_expiry_using_the_current_secret(http, clock, auth_method): + secret = ["initial-secret"] + observed = [] + + def load_secret(): + observed.append(secret[0]) + return secret[0] + + http.serve(TOKEN_URL, {"access_token": "first", "token_type": "Bearer", "expires_in": 100}, method="POST") + subject = client(client_secret=load_secret, auth_method=auth_method) + assert subject.auth_headers()["Authorization"] == "Bearer first" + clock.advance(69) + assert subject.auth_headers()["Authorization"] == "Bearer first" + assert observed == ["initial-secret"] + secret[0] = "rotated-secret" + http.serve(TOKEN_URL, {"access_token": "second", "token_type": "Bearer", "expires_in": 100}, method="POST") + clock.advance(1) + + assert subject.auth_headers()["Authorization"] == "Bearer second" + assert observed == ["initial-secret", "rotated-secret"] + + +@pytest.mark.parametrize("lifetime", [1, 30, None]) +def test_short_lived_tokens_and_tokens_without_lifetimes_are_not_cached(http, lifetime): + payload = {"access_token": "first", "token_type": "Bearer"} + if lifetime is not None: + payload["expires_in"] = lifetime + http.serve(TOKEN_URL, payload, method="POST") + subject = client() + assert subject.auth_headers()["Authorization"] == "Bearer first" + http.serve(TOKEN_URL, {**payload, "access_token": "second"}, method="POST") + + assert subject.auth_headers()["Authorization"] == "Bearer second" + assert len(http.requests) == 2 + + +@pytest.mark.parametrize( + "override", + [ + {"access_token": ""}, + {"access_token": None}, + {"access_token": "token\r\ninjected"}, + {"token_type": "DPoP"}, + {"token_type": None}, + {"expires_in": "3600"}, + {"expires_in": 0}, + {"expires_in": -1}, + {"expires_in": True}, + {"expires_in": None}, + {"expires_in": float("inf")}, + ], +) +def test_invalid_token_responses_are_rejected_without_retry(http, override): + payload = {"access_token": "token", "token_type": "Bearer", "expires_in": 3600, **override} + http.serve(TOKEN_URL, payload, method="POST") + + subject = client() + with pytest.raises(TokenExchangeError): + subject.auth_headers() + assert len(http.requests) == 1 + + +def test_resources_have_separate_token_caches_and_request_parameters(http): + http.responses[("POST", TOKEN_URL)] = deque( + [ + (200, {"access_token": "orders-token", "token_type": "Bearer", "expires_in": 3600}), + (200, {"access_token": "inventory-token", "token_type": "Bearer", "expires_in": 3600}), + ], + ) + orders = client(audience="https://orders.example.com") + inventory = client(resource="https://inventory.example.com") + + assert orders.auth_headers()["Authorization"] == "Bearer orders-token" + assert inventory.auth_headers()["Authorization"] == "Bearer inventory-token" + assert orders.auth_headers()["Authorization"] == "Bearer orders-token" + assert len(http.requests) == 2 + assert parse_qs(http.requests[0][2]["body"].decode())["audience"] == ["https://orders.example.com"] + assert parse_qs(http.requests[1][2]["body"].decode())["resource"] == ["https://inventory.example.com"] + + +@pytest.mark.parametrize("status", [400, 401, 403]) +def test_permanent_exchange_errors_are_not_retried(http, status): + http.serve( + TOKEN_URL, + {"error": "invalid_client", "error_description": "private details"}, + status=status, + method="POST", + ) + + subject = client() + with pytest.raises(TokenExchangeError): + subject.auth_headers() + assert len(http.requests) == 1 + + +def test_transient_exchange_errors_have_at_most_two_retries(http, clock, monkeypatch, auth_method): + http.serve(TOKEN_URL, b"temporarily unavailable", status=503, method="POST") + monkeypatch.setattr(time, "sleep", clock.advance) + secrets = [] + + def load_secret(): + secrets.append("secret") + return secrets[-1] + + subject = client(client_secret=load_secret, auth_method=auth_method) + with pytest.raises(TokenExchangeError): + subject.auth_headers() + assert len(http.requests) == 3 + assert len(secrets) == 3 + + +def test_transient_exchange_can_recover_within_the_same_budget(http, clock, monkeypatch): + http.responses[("POST", TOKEN_URL)] = deque( + [(429, {}), (200, {"access_token": "recovered", "token_type": "Bearer", "expires_in": 100})], + ) + monkeypatch.setattr(time, "sleep", clock.advance) + + assert client().auth_headers() == {"Authorization": "Bearer recovered"} + assert len(http.requests) == 2 + + +def test_exchange_cannot_accept_a_response_after_its_deadline(http, clock): + def slow_endpoint(): + clock.advance(4) + return {"access_token": "too-late", "token_type": "Bearer", "expires_in": 3600} + + http.serve(TOKEN_URL, slow_endpoint, method="POST") + subject = client(timeout_seconds=3) + with pytest.raises(TokenExchangeError): + subject.auth_headers() + assert len(http.requests) == 1 + + +def test_exchange_cannot_return_a_token_that_expired_during_the_request(http, clock): + def slow_endpoint(): + clock.advance(2) + return {"access_token": "already-expired", "token_type": "Bearer", "expires_in": 1} + + http.serve(TOKEN_URL, slow_endpoint, method="POST") + subject = client() + with pytest.raises(TokenExchangeError): + subject.auth_headers() + + +def test_concurrent_requests_share_one_token_exchange(http, auth_method): + entered = threading.Event() + release = threading.Event() + + def exchange(): + entered.set() + assert release.wait(2) + return {"access_token": "shared-token", "token_type": "Bearer", "expires_in": 100} + + http.serve(TOKEN_URL, exchange, method="POST") + subject = client(auth_method=auth_method) + with ThreadPoolExecutor(max_workers=8) as executor: + results = [executor.submit(subject.auth_headers) for _ in range(8)] + assert entered.wait(2) + release.set() + assert all(result.result(timeout=2) == {"Authorization": "Bearer shared-token"} for result in results) + assert len(http.requests) == 1 + + +@pytest.mark.parametrize("status", [200, 401], ids=["success", "failure"]) +def test_concurrent_callers_share_reacquisition_at_the_refresh_boundary(http, clock, monkeypatch, status): + http.serve(TOKEN_URL, {"access_token": "old", "token_type": "Bearer", "expires_in": 100}, method="POST") + subject = client() + assert subject.auth_headers() == {"Authorization": "Bearer old"} + clock.advance(69) + assert subject.auth_headers() == {"Authorization": "Bearer old"} + assert len(http.requests) == 1 + clock.advance(1) # The old token is still valid, with exactly 30 seconds remaining. + + entered = threading.Event() + release = threading.Event() + joined = threading.Event() + + def exchange(): + entered.set() + assert release.wait(5) + if status == 200: + return {"access_token": "replacement", "token_type": "Bearer", "expires_in": 100} + return {"error": "invalid_client"} + + http.serve(TOKEN_URL, exchange, status=status, method="POST") + with ThreadPoolExecutor(max_workers=3) as executor: + owner = executor.submit(subject.auth_headers) + waiters = [] + try: + assert entered.wait(5) + flight = subject._flight + assert flight is not None + wait = flight.done.wait + + def observe_wait(timeout): + joined.set() + return wait(timeout) + + monkeypatch.setattr(flight.done, "wait", observe_wait) + for _ in range(2): + joined.clear() + waiters.append(executor.submit(subject.auth_headers)) + assert joined.wait(5) + assert len(http.requests) == 2 + assert not owner.done() + assert all(not waiter.done() for waiter in waiters) + finally: + release.set() + + for result in (owner, *waiters): + if status == 200: + assert result.result(timeout=5) == {"Authorization": "Bearer replacement"} + else: + with pytest.raises(TokenExchangeError): + result.result(timeout=5) + assert len(http.requests) == 2 + + if status == 200: + assert subject.auth_headers() == {"Authorization": "Bearer replacement"} + assert len(http.requests) == 2 + else: + # A later caller must also exchange again instead of using the still-valid old token. + with pytest.raises(TokenExchangeError): + subject.auth_headers() + assert len(http.requests) == 3 + + +def test_secret_loader_errors_and_representations_are_redacted(http): + def load_secret(): + raise RuntimeError("sensitive-loader-data") + + subject = client(client_secret=load_secret) + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + assert "sensitive-loader-data" not in "".join(traceback.format_exception(error.value)) + assert repr(subject) == "" + + +@pytest.mark.parametrize( + "options", + [ + {"audience": "one", "resource": "two"}, + {"audience": " "}, + {"resource": ""}, + {"resource": 42}, + {"token_url": "http://idp.example.com/token"}, + {"token_url": "https://user:secret@idp.example.com/token"}, + {"client_id": ""}, + {"client_secret": ""}, + {"timeout_seconds": 0}, + {"timeout_seconds": float("inf")}, + {"scopes": ["scope\ninjection"]}, + ], +) +def test_invalid_client_configuration_is_rejected(options): + with pytest.raises(ValueError): + client(**options) + + +def test_request_attaches_resource_token_without_forwarding_client_credentials(http): + http.serve(TOKEN_URL, {"access_token": "resource-token", "token_type": "Bearer", "expires_in": 100}, method="POST") + http.serve("https://api.example.com/orders", {"orders": [123]}) + subject = client(audience="https://api.example.com") + + response = subject.request( + "GET", + "https://api.example.com/orders", + headers={"Accept": "application/json"}, + timeout=5, + ) + + assert response.json() == {"orders": [123]} + request = http.requests[-1][2] + assert request["headers"] == {"Accept": "application/json", "Authorization": "Bearer resource-token"} + assert request["redirect"] is False + assert request["retries"] is False + assert "test-client-secret" not in repr(subject) + assert "resource-token" not in repr(subject) + + +@pytest.mark.parametrize( + "url,options", + [ + ("http://api.example.com/orders", {}), + ("https://api.example.com/orders", {"headers": {"authorization": "other-token"}}), + ("https://api.example.com/orders", {"redirect": True}), + ("https://api.example.com/orders", {"retries": 3}), + ("https://api.example.com/orders", {"timeout": 0}), + ("https://api.example.com/orders", {"headers": [("Accept", "application/json")]}), + ], +) +def test_request_rejects_unsafe_overrides_before_acquiring_credentials(http, url, options): + subject = client() + with pytest.raises(ValueError): + subject.request("GET", url, **options) + assert http.requests == [] + + +def test_request_does_not_follow_redirects_or_retry_downstream_failures(http): + http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST") + http.serve("https://api.example.com/orders", {}, status=302) + subject = client() + + assert subject.request("GET", "https://api.example.com/orders").status == 302 + http.serve("https://api.example.com/orders", {}, status=503) + assert subject.request("GET", "https://api.example.com/orders").status == 503 + assert len(http.requests) == 3 + + +@pytest.mark.parametrize("method", [None, "", "GET /", "GET\r\nInjected"]) +def test_invalid_http_methods_are_rejected_before_loading_credentials(http, method): + calls = [] + + def load_secret(): + calls.append(True) + return "test-secret" + + subject = client(client_secret=load_secret) + with pytest.raises(ValueError, match="HTTP method"): + subject.request(method, "https://api.example.com/orders") + assert calls == [] + assert http.requests == [] + + +@pytest.mark.parametrize("secret", [None, "", 42]) +def test_invalid_secret_loader_results_are_rejected_before_sending_credentials(http, secret, auth_method): + subject = client(client_secret=lambda: secret, auth_method=auth_method) + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + assert error.value.__context__ is None + assert http.requests == [] + + +def test_retry_stops_when_the_backoff_exceeds_the_remaining_budget(http, clock, monkeypatch): + sleeps = [] + monkeypatch.setattr(time, "sleep", sleeps.append) + http.serve(TOKEN_URL, {}, status=503, method="POST") + + subject = client(timeout_seconds=0.05) + with pytest.raises(TokenExchangeError): + subject.auth_headers() + assert sleeps == [] + assert len(http.requests) == 1 + + +def test_waiting_callers_share_a_failed_exchange_and_can_recover(http, monkeypatch): + entered = threading.Event() + release = threading.Event() + joined = threading.Event() + + def exchange(): + entered.set() + assert release.wait(5) + return {"error": "invalid_client"} + + http.serve(TOKEN_URL, exchange, status=401, method="POST") + subject = client() + with ThreadPoolExecutor(max_workers=2) as executor: + owner = executor.submit(subject.auth_headers) + try: + assert entered.wait(5) + flight = subject._flight + assert flight is not None + wait = flight.done.wait + + def observe_wait(timeout): + joined.set() + return wait(timeout) + + # Keep the real Event; observe it so the provider is released only + # after the second caller has joined the active exchange. + monkeypatch.setattr(flight.done, "wait", observe_wait) + waiter = executor.submit(subject.auth_headers) + assert joined.wait(5) + finally: + release.set() + for result in (owner, waiter): + with pytest.raises(TokenExchangeError) as error: + result.result(timeout=5) + assert error.value.__context__ is None + assert len(http.requests) == 1 + + http.serve(TOKEN_URL, {"access_token": "recovered", "token_type": "Bearer", "expires_in": 100}, method="POST") + assert subject.auth_headers() == {"Authorization": "Bearer recovered"} + assert len(http.requests) == 2 + + +def test_waiting_callers_timeout_without_returning_the_late_token(http): + entered = threading.Event() + release = threading.Event() + + def exchange(): + entered.set() + assert release.wait(5) + return {"access_token": "too-late", "token_type": "Bearer", "expires_in": 100} + + http.serve(TOKEN_URL, exchange, method="POST") + subject = client(timeout_seconds=0.1) + with ThreadPoolExecutor(max_workers=1) as executor: + owner = executor.submit(subject.auth_headers) + try: + assert entered.wait(5) + with pytest.raises(TokenExchangeError): + subject.auth_headers() + assert not owner.done() + assert len(http.requests) == 1 + finally: + release.set() + with pytest.raises(TokenExchangeError): + owner.result(timeout=5) + + http.serve(TOKEN_URL, {"access_token": "recovered", "token_type": "Bearer", "expires_in": 100}, method="POST") + assert subject.auth_headers() == {"Authorization": "Bearer recovered"} + assert len(http.requests) == 2 + + +def test_configuration_and_returned_headers_do_not_mutate_the_token_cache(http): + scopes = ["orders:read"] + subject = client(scopes=scopes) + scopes.append("orders:write") + http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST") + + headers = subject.auth_headers() + headers["Authorization"] = "Bearer replacement" + assert subject.auth_headers() == {"Authorization": "Bearer token"} + assert parse_qs(http.requests[0][2]["body"].decode())["scope"] == ["orders:read"] + assert len(http.requests) == 1 + + +def test_unselected_resource_and_scopes_are_not_added_to_the_exchange(http): + subject = client(scopes=[]) + http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST") + + subject.auth_headers() + assert parse_qs(http.requests[0][2]["body"].decode()) == {"grant_type": ["client_credentials"]} + + +@pytest.mark.parametrize("lifetime", [1, 30, None]) +def test_concurrent_callers_share_an_uncacheable_token(http, monkeypatch, lifetime): + entered = threading.Event() + release = threading.Event() + joined = threading.Event() + payload = {"access_token": "shared", "token_type": "Bearer"} + if lifetime is not None: + payload["expires_in"] = lifetime + + def exchange(): + entered.set() + assert release.wait(5) + return payload + + subject = client() + http.serve(TOKEN_URL, exchange, method="POST") + with ThreadPoolExecutor(max_workers=2) as executor: + owner = executor.submit(subject.auth_headers) + try: + assert entered.wait(5) + flight = subject._flight + assert flight is not None + wait = flight.done.wait + + def observe_wait(timeout): + joined.set() + return wait(timeout) + + monkeypatch.setattr(flight.done, "wait", observe_wait) + waiter = executor.submit(subject.auth_headers) + assert joined.wait(5) + finally: + release.set() + assert owner.result(timeout=5) == {"Authorization": "Bearer shared"} + assert waiter.result(timeout=5) == {"Authorization": "Bearer shared"} + assert len(http.requests) == 1 + subject.auth_headers() + assert len(http.requests) == 2 + + +def test_invalid_unicode_from_a_secret_loader_is_sanitized(http, auth_method): + subject = client(client_secret=lambda: "private-secret-\ud800", auth_method=auth_method) + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + assert error.value.__context__ is None + assert "private-secret" not in "".join(traceback.format_exception(error.value)) + assert http.requests == [] + + +@pytest.mark.parametrize("field", ["client_id", "audience", "resource"]) +def test_invalid_configuration_encoding_is_rejected(field): + options = {field: "\ud800"} + with pytest.raises(ValueError, match="UTF-8"): + client(**options) + + +@pytest.mark.parametrize( + "resource", + [ + "inventory", + "/inventory", + "//inventory.example.com", + "https://inventory.example.com/#fragment", + "https://inventory.example.com/#", + "urn:example:inventory#fragment", + " https://inventory.example.com", + "https://inventory.example.com/\nstock", + "https://inventory.example.com/\x00stock", + "https://inventory.example.com/stock item", + "https://inventory.example.com/%invalid", + "https://[invalid", + "1inventory:stock", + "urn%3Ainventory", + "urn:inventory%", + "urn:inventory%2", + "urn:inventory%2G", + "urn:inventory%G2", + "urn:inventory%%20", + "urn:inventorý", + ], +) +def test_resource_must_be_an_absolute_uri_without_a_fragment(resource): + with pytest.raises(ValueError, match="absolute URI without a fragment"): + client(resource=resource) + + +@pytest.mark.parametrize( + "selection", + [ + {"resource": "urn:example:inventory"}, + {"resource": "https://inventory.example.com/stock?region=eu&category=%23parts"}, + {"resource": "http://inventory.example.com"}, + {"resource": "URN:example:inventory"}, + {"resource": "inventory+v1.2-test:stock%2fitems%20eu?category=%23"}, + {"resource": "https://[::1]/caf%C3%A9?encoded=%00%ff"}, + {"audience": "inventory"}, + ], +) +def test_valid_resource_identifiers_and_provider_audiences_are_preserved(http, selection): + http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST") + subject = client(**selection) + + subject.auth_headers() + + name, value = next(iter(selection.items())) + fields = parse_qs(http.requests[0][2]["body"].decode()) + assert fields[name] == [value] + other = "audience" if name == "resource" else "resource" + assert other not in fields + + +@pytest.mark.parametrize("suffix", ["%", "%4", "%4G", "#"]) +def test_long_resources_with_invalid_suffixes_are_rejected_before_loading_credentials(http, suffix): + secrets = [] + + def load_secret(): + secrets.append("test-secret") + return secrets[-1] + + resource = "urn:inventory:" + "a%41" * 25_000 + suffix + with pytest.raises(ValueError, match="absolute URI without a fragment"): + client(resource=resource, client_secret=load_secret) + + assert secrets == [] + assert http.requests == [] + + +def test_secret_lookup_does_not_expire_a_new_short_lived_token(http, clock, auth_method): + def load_secret(): + clock.advance(2) + return "test-secret" + + http.serve(TOKEN_URL, {"access_token": "fresh", "token_type": "Bearer", "expires_in": 1}, method="POST") + subject = client(client_secret=load_secret, timeout_seconds=3, auth_method=auth_method) + + assert subject.auth_headers() == {"Authorization": "Bearer fresh"} + assert subject.auth_headers() == {"Authorization": "Bearer fresh"} + assert len(http.requests) == 2 + + +def test_secret_lookup_does_not_move_the_cached_tokens_refresh_boundary(http, clock, auth_method): + def load_secret(): + clock.advance(2) + return "test-secret" + + http.serve(TOKEN_URL, {"access_token": "first", "token_type": "Bearer", "expires_in": 100}, method="POST") + subject = client(client_secret=load_secret, auth_method=auth_method) + assert subject.auth_headers() == {"Authorization": "Bearer first"} + clock.advance(69) + assert subject.auth_headers() == {"Authorization": "Bearer first"} + assert len(http.requests) == 1 + + http.serve(TOKEN_URL, {"access_token": "second", "token_type": "Bearer", "expires_in": 100}, method="POST") + clock.advance(1) + assert subject.auth_headers() == {"Authorization": "Bearer second"} + assert len(http.requests) == 2 + + +def test_secret_lookup_still_consumes_the_acquisition_budget(http, clock, auth_method): + def load_secret(): + clock.advance(4) + return "test-secret" + + subject = client(client_secret=load_secret, timeout_seconds=3, auth_method=auth_method) + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + + assert error.value.retryable + assert http.requests == [] + + +@pytest.mark.parametrize( + "name", + [ + "Authorization", + "aUtHoRiZaTiOn", + "Authorization ", + "Authorization\t", + " Authorization", + "", + ":", + "X:Trace", + "X Trace", + "X\tTrace", + "X/Trace", + "X(Trace)", + "X\x00Trace", + "X\x7fTrace", + "X-Ünicode", + "X-Trace\r\nInjected", + ], +) +def test_invalid_header_names_are_rejected_before_token_acquisition(http, name): + secrets = [] + + def load_secret(): + secrets.append("test-secret") + return secrets[-1] + + http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST") + http.serve("https://api.example.com/orders", {"orders": []}) + subject = client(client_secret=load_secret) + + with pytest.raises(ValueError, match="Request headers"): + subject.request("GET", "https://api.example.com/orders", headers={name: "test-value"}) + + assert secrets == [] + assert http.requests == [] + + +@pytest.mark.parametrize("value", ["\x00", "\x01", "\x1f", "\x7f", "東京", "\ud800"]) +def test_invalid_header_values_are_rejected_before_token_acquisition(http, value): + secrets = [] + + def load_secret(): + secrets.append("test-secret") + return secrets[-1] + + http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST") + http.serve("https://api.example.com/orders", {"orders": []}) + subject = client(client_secret=load_secret) + + with pytest.raises(ValueError, match="Request headers"): + subject.request("GET", "https://api.example.com/orders", headers={"X-Trace": f"trace{value}value"}) + + assert secrets == [] + assert http.requests == [] + + +@pytest.mark.parametrize("value", ["", "trace-id", "trace\tvalue", "\x80", "\xff", "caf\xe9"]) +def test_valid_header_names_and_values_are_preserved(http, value): + http.serve(TOKEN_URL, {"access_token": "token", "token_type": "Bearer", "expires_in": 100}, method="POST") + http.serve("https://api.example.com/orders", {"orders": []}) + subject = client() + headers = {"X-Trace!#$%&'*+.^_`|~09": value} + + response = subject.request("GET", "https://api.example.com/orders", headers=headers) + + assert response.status == 200 + assert http.requests[-1][2]["headers"] == {**headers, "Authorization": "Bearer token"} diff --git a/tests/functional/auth_alpha/oauth2/test_errors.py b/tests/functional/auth_alpha/oauth2/test_errors.py new file mode 100644 index 00000000000..62cf783cc59 --- /dev/null +++ b/tests/functional/auth_alpha/oauth2/test_errors.py @@ -0,0 +1,109 @@ +import io +import json +import traceback +from uuid import uuid4 + +import pytest +import urllib3 + +from aws_lambda_powertools import Logger +from aws_lambda_powertools.utilities.auth_alpha import AuthFailureReason, OAuth2Client +from aws_lambda_powertools.utilities.auth_alpha.exceptions import AuthError +from aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions import DownstreamRequestError, TokenExchangeError + +TOKEN_URL = "https://idp.example.com/token" +RESOURCE_URL = "https://api.example.com/orders" +PRIVATE_DATA = "test-only-sensitive-provider-data" + + +@pytest.mark.parametrize("operation", ["auth_headers", "request"]) +@pytest.mark.parametrize("failure", ["secret", "transport", "json", "expires_in", "downstream"]) +@pytest.mark.parametrize("auth_method", ["client_secret_basic", "client_secret_post"]) +def test_errors_never_expose_credentials_or_active_exception_chains(http, operation, failure, monkeypatch, auth_method): + monkeypatch.setattr("time.sleep", lambda seconds: None) + + def load_secret(): + if failure == "secret": + raise RuntimeError(PRIVATE_DATA) + return PRIVATE_DATA + + subject = OAuth2Client( + token_url=TOKEN_URL, + client_id="orders", + client_secret=load_secret, + auth_method=auth_method, + ) + payload = {"access_token": "test-token", "token_type": "Bearer", "expires_in": 600} + if failure == "transport": + payload = urllib3.exceptions.SSLError(PRIVATE_DATA) + elif failure == "json": + payload = PRIVATE_DATA.encode() + elif failure == "expires_in": + payload["expires_in"] = PRIVATE_DATA + elif operation == "auth_headers" and failure == "downstream": + # This operation has no downstream request, so exercise a bad token instead. + payload["access_token"] = PRIVATE_DATA + "\r\ninvalid" + http.serve(TOKEN_URL, payload, method="POST") + http.serve(RESOURCE_URL, urllib3.exceptions.SSLError(PRIVATE_DATA)) + stream = io.StringIO() + logger = Logger(service=f"oauth-error-test-{uuid4()}", stream=stream) + expected = DownstreamRequestError if operation == "request" and failure == "downstream" else TokenExchangeError + invoke = subject.request if operation == "request" else subject.auth_headers + args = ("GET", RESOURCE_URL) if operation == "request" else () + + try: + raise LookupError(PRIVATE_DATA) + except LookupError: + with pytest.raises(expected) as captured: + invoke(*args) + error = captured.value + logger.exception("Authentication failed", exc_info=(type(error), error, error.__traceback__)) + assert isinstance(error, AuthError) + assert error.__context__ is None + assert error.__cause__ is None + assert PRIVATE_DATA not in str(error) + assert PRIVATE_DATA not in repr(error) + assert PRIVATE_DATA not in "".join(traceback.format_exception(error)) + assert PRIVATE_DATA not in stream.getvalue() + assert json.loads(stream.getvalue())["exception_name"] == expected.__name__ + + +@pytest.mark.parametrize( + "status,retryable,attempts", + [(400, False, 1), (401, False, 1), (429, True, 3), (503, True, 3)], +) +@pytest.mark.parametrize("auth_method", ["client_secret_basic", "client_secret_post"]) +def test_exchange_errors_expose_fixed_reason_and_retryability( + http, + clock, + monkeypatch, + status, + retryable, + attempts, + auth_method, +): + monkeypatch.setattr("time.sleep", clock.advance) + subject = OAuth2Client( + token_url=TOKEN_URL, + client_id="orders", + client_secret="test-secret", + auth_method=auth_method, + ) + http.serve(TOKEN_URL, {"error_description": PRIVATE_DATA}, method="POST", status=status) + + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + assert error.value.reason is AuthFailureReason.TOKEN_EXCHANGE_FAILED + assert error.value.retryable is retryable + assert len(http.requests) == attempts + + +def test_outbound_and_jwt_errors_share_the_public_base(): + from aws_lambda_powertools.utilities.auth_alpha.jwt.exceptions import AuthError as JWTAuthError + from aws_lambda_powertools.utilities.auth_alpha.jwt.exceptions import AuthFailureReason as JWTFailureReason + + assert JWTAuthError is AuthError + assert JWTFailureReason is AuthFailureReason + assert issubclass(TokenExchangeError, AuthError) + assert DownstreamRequestError().reason is AuthFailureReason.DOWNSTREAM_REQUEST_FAILED + assert not DownstreamRequestError().retryable diff --git a/tests/functional/auth_alpha/oauth2/test_imports.py b/tests/functional/auth_alpha/oauth2/test_imports.py new file mode 100644 index 00000000000..7ea0f42adcc --- /dev/null +++ b/tests/functional/auth_alpha/oauth2/test_imports.py @@ -0,0 +1,44 @@ +import os +import subprocess +import sys +from pathlib import Path + + +def test_oauth_client_does_not_import_jwt_or_cryptography(): + root = Path(__file__).parents[4] + probe = """ +import importlib.abc +import sys + +class BlockJWT(importlib.abc.MetaPathFinder): + def find_spec(self, fullname, path=None, target=None): + if fullname.split(".")[0] in {"jwt", "cryptography"}: + raise AssertionError("OAuth imported JWT dependencies") + +sys.meta_path.insert(0, BlockJWT()) +import aws_lambda_powertools.utilities.auth_alpha as auth +import aws_lambda_powertools.utilities.auth_alpha.oauth2 as oauth +assert "urllib3" not in sys.modules +assert "OAuth2Client" in dir(auth) +assert "OAuth2Client" in dir(oauth) +assert auth.OAuth2Client is oauth.OAuth2Client +for auth_method in ("client_secret_basic", "client_secret_post"): + client = oauth.OAuth2Client( + token_url="https://idp.example.com/token", + client_id="orders", + client_secret="test-secret", + auth_method=auth_method, + ) + assert repr(client) == "" +try: + oauth.unknown_attribute +except AttributeError: + pass +else: + raise AssertionError("Unexpected module attribute") +assert not {"jwt", "cryptography"} & sys.modules.keys() +""" + env = os.environ.copy() + env["PYTHONPATH"] = os.pathsep.join((str(root), env.get("PYTHONPATH", ""))) + result = subprocess.run([sys.executable, "-c", probe], capture_output=True, text=True, env=env, check=False) + assert result.returncode == 0, result.stderr diff --git a/tests/integration/auth_alpha/jwt/conftest.py b/tests/integration/auth_alpha/conftest.py similarity index 58% rename from tests/integration/auth_alpha/jwt/conftest.py rename to tests/integration/auth_alpha/conftest.py index 78304d5c565..d95a5a74a11 100644 --- a/tests/integration/auth_alpha/jwt/conftest.py +++ b/tests/integration/auth_alpha/conftest.py @@ -1,12 +1,15 @@ """A local TLS endpoint exercising the production transport without HTTP mocks.""" +from __future__ import annotations + import ipaddress import json import ssl import threading -from dataclasses import dataclass, field +from dataclasses import dataclass, field, replace from datetime import datetime, timedelta, timezone from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import TYPE_CHECKING import pytest from cryptography import x509 @@ -14,6 +17,9 @@ from cryptography.hazmat.primitives.asymmetric import rsa from cryptography.x509.oid import NameOID +if TYPE_CHECKING: + from collections.abc import Callable + @dataclass class Reply: @@ -22,6 +28,50 @@ class Reply: headers: dict = field(default_factory=dict) interval: float = 0 stall: bool = False + header_interval: float = 0 + chunked: bool = False + chunk_size_interval: float = 0 + trailer_interval: float = 0 + responder: Callable[[str, dict, bytes], tuple[int, dict]] | None = None + + +def _write_bytes(stream, payload, stop, interval=0): + if not interval: + stream.write(payload) + return True + for value in payload: + if stop.wait(interval): + return False + stream.write(bytes([value])) + stream.flush() + return True + + +def _write_body(stream, reply, stop): + if reply.stall: + stop.wait(5) + return + parts = [(reply.body, reply.interval)] + if reply.chunked: + parts = [] + if reply.body: + parts.extend( + [ + (f"{len(reply.body):x};padding=".encode() + b"x" * 32 + b"\r\n", reply.chunk_size_interval), + (reply.body, reply.interval), + (b"\r\n", 0), + ], + ) + parts.extend( + [ + (b"0\r\n", 0), + (b"X-Trailer: " + b"x" * 32 + b"\r\n", reply.trailer_interval), + (b"\r\n", 0), + ], + ) + for payload, interval in parts: + if not _write_bytes(stream, payload, stop, interval): + return class LocalHTTPS: @@ -31,9 +81,34 @@ def __init__(self): self.stop = threading.Event() self.url = "" - def serve(self, path, payload, *, status=200, headers=None, interval=0, stall=False): + def serve( + self, + path, + payload, + *, + status=200, + headers=None, + interval=0, + stall=False, + header_interval=0, + chunked=False, + chunk_size_interval=0, + trailer_interval=0, + responder=None, + ): body = payload if isinstance(payload, bytes) else json.dumps(payload).encode() - self.routes[path] = Reply(body, status, headers or {}, interval, stall) + self.routes[path] = Reply( + body=body, + status=status, + headers=headers or {}, + interval=interval, + stall=stall, + header_interval=header_interval, + chunked=chunked, + chunk_size_interval=chunk_size_interval, + trailer_interval=trailer_interval, + responder=responder, + ) @pytest.fixture(scope="session") @@ -90,28 +165,38 @@ def do_GET(self): # noqa: N802 def do_POST(self): # noqa: N802 self.respond() + def do_HEAD(self): # noqa: N802 + self.respond() + def respond(self): body = self.rfile.read(int(self.headers.get("Content-Length", 0))) endpoint.requests.append((self.command, self.path, dict(self.headers), body)) reply = endpoint.routes.get(self.path, Reply(b"{}", status=404)) + if reply.responder is not None: + status, payload = reply.responder(self.command, dict(self.headers), body) + reply = replace(reply, status=status, body=json.dumps(payload).encode()) self.send_response(reply.status) self.send_header("Content-Type", "application/json") - self.send_header("Content-Length", str(len(reply.body))) + if reply.chunked: + self.send_header("Transfer-Encoding", "chunked") + else: + self.send_header("Content-Length", str(len(reply.body))) self.send_header("Connection", "keep-alive" if request.param else "close") for name, value in reply.headers.items(): self.send_header(name, value) - self.end_headers() try: - if reply.stall: - endpoint.stop.wait(5) - elif reply.interval: - for value in reply.body: - if endpoint.stop.wait(reply.interval): - break - self.wfile.write(bytes([value])) - self.wfile.flush() - else: - self.wfile.write(reply.body) + if reply.header_interval: + self.flush_headers() + if not _write_bytes( + self.wfile, + b"X-Slow: " + b"x" * 32 + b"\r\n", + endpoint.stop, + reply.header_interval, + ): + return + self.end_headers() + if self.command != "HEAD": + _write_body(self.wfile, reply, endpoint.stop) except (OSError, ssl.SSLError): # Timeout and oversized-body tests deliberately close early. pass diff --git a/tests/integration/auth_alpha/jwt/test_https.py b/tests/integration/auth_alpha/jwt/test_https.py index bece05d6c65..13dde0f6a38 100644 --- a/tests/integration/auth_alpha/jwt/test_https.py +++ b/tests/integration/auth_alpha/jwt/test_https.py @@ -1,4 +1,5 @@ import time +import traceback import jwt import pytest @@ -70,3 +71,19 @@ def test_key_endpoint_failures_are_bounded_and_do_not_follow_redirects(https_ser assert time.monotonic() - started < 1 assert error.value.__context__ is None assert [request[1] for request in https_server.requests] == ["/keys"] + + +def test_malformed_jwks_headers_fail_without_logging_provider_data(https_server, caplog): + private_data = "local-test-private-provider-data" + https_server.serve("/keys", {"keys": []}, headers={"Broken header": private_data}) + subject = verifier(https_server, jwks_uri=https_server.url + "/keys") + + with pytest.raises(JWKSFetchError) as error: + subject.prefetch() + + assert error.value.__context__ is None + assert error.value.__cause__ is None + assert private_data not in "".join(traceback.format_exception(error.value)) + assert private_data not in caplog.text + assert not [record for record in caplog.records if record.name == "urllib3.connection"] + assert [request[1] for request in https_server.requests] == ["/keys"] diff --git a/tests/integration/auth_alpha/oauth2/__init__.py b/tests/integration/auth_alpha/oauth2/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/auth_alpha/oauth2/test_https.py b/tests/integration/auth_alpha/oauth2/test_https.py new file mode 100644 index 00000000000..01a9f147d10 --- /dev/null +++ b/tests/integration/auth_alpha/oauth2/test_https.py @@ -0,0 +1,286 @@ +import base64 +import threading +import time +import traceback +from concurrent.futures import ThreadPoolExecutor +from urllib.parse import parse_qs + +import pytest + +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions import DownstreamRequestError, TokenExchangeError + +TOKEN_RESPONSE = {"access_token": "local-test-token", "token_type": "Bearer", "expires_in": 600} + + +@pytest.fixture(params=["client_secret_basic", "client_secret_post"]) +def auth_method(request): + return request.param + + +def client(endpoint, **options): + return OAuth2Client( + token_url=endpoint.url + "/token", + client_id="orders", + client_secret="test-only-secret", + scopes=["inventory:read"], + **options, + ) + + +@pytest.mark.parametrize("chunked", [False, True]) +def test_token_exchange_and_authenticated_request_over_trusted_tls(https_server, chunked): + https_server.serve("/token", TOKEN_RESPONSE, chunked=chunked) + https_server.serve("/inventory", {"items": [123]}, chunked=chunked) + subject = client(https_server) + + response = subject.request("GET", https_server.url + "/inventory") + + assert response.status == 200 + assert response.json() == {"items": [123]} + assert subject.auth_headers() == {"Authorization": "Bearer local-test-token"} + exchange, resource = https_server.requests + assert exchange[:2] == ("POST", "/token") + assert base64.b64decode(exchange[2]["Authorization"].removeprefix("Basic ")).decode() == "orders:test-only-secret" + assert parse_qs(exchange[3].decode()) == {"grant_type": ["client_credentials"], "scope": ["inventory:read"]} + assert resource[2]["Authorization"] == "Bearer local-test-token" + assert "test-only-secret" not in str(resource) + + +def test_downstream_redirects_are_returned_without_forwarding_bearer_tokens(https_server): + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", {}, status=307, headers={"Location": https_server.url + "/other"}) + https_server.serve("/other", {}) + + response = client(https_server).request("GET", https_server.url + "/inventory") + + assert response.status == 307 + assert [request[1] for request in https_server.requests] == ["/token", "/inventory"] + + +def test_downstream_failures_are_not_retried(https_server): + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", {}, status=503) + assert client(https_server).request("POST", https_server.url + "/inventory").status == 503 + assert [request[1] for request in https_server.requests] == ["/token", "/inventory"] + + +def test_downstream_timeout_has_a_separate_budget_and_a_sanitized_error(https_server): + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", {"items": []}, stall=True) + subject = client(https_server, timeout_seconds=3) + started = time.monotonic() + + with pytest.raises(DownstreamRequestError) as error: + subject.request("GET", https_server.url + "/inventory", timeout=0.2) + + assert time.monotonic() - started < 1 + assert error.value.__context__ is None + assert [request[1] for request in https_server.requests] == ["/token", "/inventory"] + + +def test_untrusted_tls_never_sends_client_credentials(https_server, monkeypatch, auth_method): + monkeypatch.delenv("SSL_CERT_FILE") + https_server.serve("/token", TOKEN_RESPONSE) + subject = client(https_server, timeout_seconds=0.15, auth_method=auth_method) + + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + assert error.value.__context__ is None + assert https_server.requests == [] + + +@pytest.mark.parametrize("failure", ["oversized", "redirect", "stall", "trickle"]) +def test_exchange_failures_are_bounded_without_forwarding_credentials(https_server, failure, auth_method): + if failure == "oversized": + https_server.serve("/token", b'{"padding":"' + b"x" * (1024 * 1024) + b'"}') + elif failure == "redirect": + https_server.serve("/token", {}, status=307, headers={"Location": https_server.url + "/redirected"}) + https_server.serve("/redirected", TOKEN_RESPONSE) + else: + https_server.serve( + "/token", + TOKEN_RESPONSE, + stall=failure == "stall", + interval=0.04 if failure == "trickle" else 0, + ) + subject = client(https_server, timeout_seconds=0.2, auth_method=auth_method) + started = time.monotonic() + + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + assert time.monotonic() - started < 1 + assert error.value.__context__ is None + assert [request[1] for request in https_server.requests] == ["/token"] + + +def test_slow_downstream_body_cannot_extend_the_timeout(https_server): + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", {"items": list(range(100))}, interval=0.04) + subject = client(https_server) + started = time.monotonic() + + with pytest.raises(DownstreamRequestError): + subject.request("GET", https_server.url + "/inventory", timeout=0.2) + assert time.monotonic() - started < 1 + assert [request[1] for request in https_server.requests] == ["/token", "/inventory"] + + +@pytest.fixture(params=["header", "chunk_size", "trailer"]) +def slow_framing(request): + return {f"{request.param}_interval": 0.04, "chunked": request.param != "header"} + + +def test_slow_token_framing_cannot_extend_the_acquisition_timeout(https_server, slow_framing): + https_server.serve("/token", TOKEN_RESPONSE, **slow_framing) + subject = client(https_server, timeout_seconds=0.2) + started = time.monotonic() + + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + + assert time.monotonic() - started < 1 + assert error.value.retryable + assert error.value.__context__ is None + assert error.value.__cause__ is None + assert [request[1] for request in https_server.requests] == ["/token"] + + https_server.serve("/token", TOKEN_RESPONSE) + assert subject.auth_headers() == {"Authorization": "Bearer local-test-token"} + assert [request[1] for request in https_server.requests] == ["/token", "/token"] + + +def test_slow_downstream_framing_cannot_extend_the_request_timeout(https_server, slow_framing): + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", {"items": [123]}, **slow_framing) + subject = client(https_server) + subject.auth_headers() + started = time.monotonic() + + with pytest.raises(DownstreamRequestError) as error: + subject.request("GET", https_server.url + "/inventory", timeout=0.2) + + assert time.monotonic() - started < 1 + assert not error.value.retryable + assert error.value.__context__ is None + assert error.value.__cause__ is None + assert [request[1] for request in https_server.requests] == ["/token", "/inventory"] + + https_server.serve("/inventory", {"items": [123]}) + assert subject.request("GET", https_server.url + "/inventory").json() == {"items": [123]} + assert [request[1] for request in https_server.requests] == ["/token", "/inventory", "/inventory"] + + +def test_concurrent_downstream_requests_keep_separate_deadlines(https_server): + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", {"items": [123]}, header_interval=0.01) + subject = client(https_server) + subject.auth_headers() + ready = threading.Barrier(2) + + def request(timeout): + ready.wait(timeout=2) + return subject.request("GET", https_server.url + "/inventory", timeout=timeout) + + with ThreadPoolExecutor(max_workers=2) as executor: + short = executor.submit(request, 0.15) + long = executor.submit(request, 2) + with pytest.raises(DownstreamRequestError): + short.result(timeout=3) + assert long.result(timeout=3).json() == {"items": [123]} + + assert [request[1] for request in https_server.requests] == ["/token", "/inventory", "/inventory"] + + +def test_body_reads_use_the_budget_remaining_after_headers(https_server): + https_server.serve("/token", TOKEN_RESPONSE) + # Header receipt consumes about 0.2s; the trailer must use what remains, + # rather than starting another downstream timeout after the headers. + https_server.serve( + "/inventory", + {"items": [123]}, + header_interval=0.005, + chunked=True, + trailer_interval=0.04, + ) + subject = client(https_server) + subject.auth_headers() + started = time.monotonic() + + with pytest.raises(DownstreamRequestError): + subject.request("GET", https_server.url + "/inventory", timeout=0.4) + + assert time.monotonic() - started < 0.6 + assert [request[1] for request in https_server.requests] == ["/token", "/inventory"] + + +@pytest.mark.parametrize("chunked", [False, True]) +def test_downstream_gzip_response_remains_readable(https_server, chunked): + import gzip + + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve( + "/inventory", + gzip.compress(b'{"items":[123]}'), + headers={"Content-Encoding": "gzip"}, + chunked=chunked, + ) + subject = client(https_server) + + response = subject.request("GET", https_server.url + "/inventory") + assert response.json() == {"items": [123]} + + +@pytest.mark.parametrize(("method", "status"), [("GET", 200), ("GET", 204), ("GET", 304), ("HEAD", 200)]) +def test_empty_downstream_responses_preserve_bytes_and_status(https_server, method, status): + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", b"not sent for HEAD" if method == "HEAD" else b"", status=status) + subject = client(https_server) + + response = subject.request(method, https_server.url + "/inventory") + + assert response.status == status + assert response.data == b"" + assert response.data.decode() == "" + assert [request[1] for request in https_server.requests] == ["/token", "/inventory"] + + +@pytest.mark.parametrize("endpoint", ["token", "inventory"]) +def test_malformed_response_headers_fail_without_logging_credentials(https_server, caplog, endpoint, auth_method): + private_data = "Bearer local-test-private-token" + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", {"items": [123]}) + payload = TOKEN_RESPONSE if endpoint == "token" else {"items": [123]} + https_server.serve(f"/{endpoint}", payload, headers={"Broken header": private_data}) + subject = client(https_server, auth_method=auth_method) + expected = TokenExchangeError if endpoint == "token" else DownstreamRequestError + + with pytest.raises(expected) as error: + subject.request("GET", https_server.url + "/inventory") + + assert not error.value.retryable + assert error.value.__context__ is None + assert error.value.__cause__ is None + assert private_data not in "".join(traceback.format_exception(error.value)) + assert private_data not in caplog.text + assert not [record for record in caplog.records if record.name == "urllib3.connection"] + expected_paths = ["/token"] if endpoint == "token" else ["/token", "/inventory"] + assert [request[1] for request in https_server.requests] == expected_paths + + https_server.serve(f"/{endpoint}", payload) + assert subject.request("GET", https_server.url + "/inventory").json() == {"items": [123]} + + +def test_valid_extended_header_values_are_sent_over_tls(https_server): + https_server.serve("/token", TOKEN_RESPONSE) + https_server.serve("/inventory", {"items": [123]}) + value = "caf\xe9\t\x80\xff" + + response = client(https_server).request( + "GET", + https_server.url + "/inventory", + headers={"X-Trace": value}, + ) + + assert response.status == 200 + assert https_server.requests[-1][2]["X-Trace"] == value diff --git a/tests/integration/auth_alpha/oauth2/test_provider_contracts.py b/tests/integration/auth_alpha/oauth2/test_provider_contracts.py new file mode 100644 index 00000000000..640292f5a84 --- /dev/null +++ b/tests/integration/auth_alpha/oauth2/test_provider_contracts.py @@ -0,0 +1,295 @@ +"""Documented provider contracts simulated over local TLS, not live provider tests. + +Profiles cover client_secret_basic and client_secret_post with client_credentials, +given the corresponding server-side permissions and application configuration. +Tokens and credentials are synthetic. These tests do not validate tenant +configuration or provider availability. + +Auth0 (select Client Secret (Basic/Post) and authorize the M2M app for the API): +https://auth0.com/docs/get-started/applications/credentials +https://auth0.com/docs/get-started/authentication-and-authorization-flow/client-credentials-flow/call-your-api-using-the-client-credentials-flow +Entra ID (application permissions, /.default scope, Basic authentication supported): +https://learn.microsoft.com/en-us/entra/identity-platform/v2-oauth2-client-creds-grant-flow +Okta (custom authorization server for your API, not Okta management API scopes): +https://developer.okta.com/docs/guides/implement-grant-type/clientcreds/main/ +https://developer.okta.com/docs/api/openapi/okta-oauth/guides/client-auth/ +https://developer.okta.com/docs/guides/implement-oauth-for-okta-serviceapp/main/ +Keycloak (client authentication and service account roles enabled): +https://www.keycloak.org/docs/latest/server_admin/index.html#_service_accounts +""" + +import base64 +import threading +import time +import traceback +from collections import Counter +from concurrent.futures import ThreadPoolExecutor +from urllib.parse import parse_qs + +import pytest + +from aws_lambda_powertools.utilities.auth_alpha import OAuth2Client +from aws_lambda_powertools.utilities.auth_alpha.oauth2.exceptions import TokenExchangeError + +PROFILES = { + "auth0": { + "path": "/oauth/token", + "options": {"audience": "https://inventory.example.com"}, + "form": {"grant_type": ["client_credentials"], "audience": ["https://inventory.example.com"]}, + "response": {"token_type": "Bearer", "expires_in": 86400}, + "error_status": 403, + "error": {"error": "access_denied", "error_description": "Unauthorized: private-auth0-detail"}, + }, + "entra": { + "path": "/test-tenant/oauth2/v2.0/token", + "options": {"scopes": ["https://graph.microsoft.com/.default"]}, + "form": {"grant_type": ["client_credentials"], "scope": ["https://graph.microsoft.com/.default"]}, + "response": {"token_type": "Bearer", "expires_in": 3599, "ext_expires_in": 3599}, + "error_status": 401, + "error": { + "error": "invalid_client", + "error_description": "AADSTS7000215: private-entra-detail", + "error_codes": [7000215], + "trace_id": "test-trace-id", + "correlation_id": "test-correlation-id", + }, + }, + "okta": { + "path": "/oauth2/default/v1/token", + "options": {"scopes": ["inventory.read", "inventory.write"]}, + "form": {"grant_type": ["client_credentials"], "scope": ["inventory.read inventory.write"]}, + "response": {"token_type": "Bearer", "expires_in": 3600, "scope": "inventory.read inventory.write"}, + "error_status": 401, + "error": {"error": "invalid_client", "error_description": "private-okta-detail"}, + }, + "keycloak": { + "path": "/realms/inventory/protocol/openid-connect/token", + "options": {}, + "form": {"grant_type": ["client_credentials"]}, + "response": { + "token_type": "Bearer", + "expires_in": 60, + "scope": "email profile", + "refresh_expires_in": 0, + "not-before-policy": 0, + }, + "error_status": 401, + "error": {"error": "invalid_client", "error_description": "private-keycloak-detail"}, + }, +} + + +class Provider: + """Reject unexpected credentials or form fields before issuing a fake token.""" + + def __init__(self, endpoint, name, auth_method="client_secret_basic"): + self.name = name + self.auth_method = auth_method + self.profile = PROFILES[name] + self.url = endpoint.url + self.profile["path"] + # Reuse the ID across providers to expose accidental cache sharing. + self.client_id = "shared-client" + self.secret = f"{name}-test-secret" + self.generation = 1 + self.forced_error = None + self.accepted = 0 + endpoint.serve(self.profile["path"], {}, responder=self.respond) + + @property + def token(self): + return f"{self.name}-token-{self.generation}" + + def client(self, **overrides): + return OAuth2Client( + **{ + "token_url": self.url, + "client_id": self.client_id, + "client_secret": self.secret, + "auth_method": self.auth_method, + **self.profile["options"], + **overrides, + }, + ) + + def respond(self, method, headers, body): + fields = parse_qs(body.decode(), keep_blank_values=True) + if self.auth_method == "client_secret_basic": + expected_basic = base64.b64encode(f"{self.client_id}:{self.secret}".encode()).decode() + if headers.get("Authorization") != f"Basic {expected_basic}": + return self.profile["error_status"], self.profile["error"] + elif ( + "Authorization" in headers + or fields.pop("client_id", None) != [self.client_id] + or fields.pop("client_secret", None) != [self.secret] + ): + return self.profile["error_status"], self.profile["error"] + if ( + method != "POST" + or headers.get("Content-Type") != "application/x-www-form-urlencoded" + or fields != self.profile["form"] + ): + return 400, {"error": "invalid_request", "error_description": "Unexpected token request"} + if self.forced_error is not None: + return self.forced_error + self.accepted += 1 + return 200, {"access_token": self.token, **self.profile["response"]} + + +@pytest.fixture(params=["client_secret_basic", "client_secret_post"]) +def auth_method(request): + return request.param + + +@pytest.fixture(params=PROFILES) +def provider(https_server, request, auth_method): + return Provider(https_server, request.param, auth_method=auth_method) + + +def test_provider_contract_acquires_caches_and_uses_token(provider, https_server): + https_server.serve("/inventory", {"items": [123]}) + subject = provider.client() + + assert subject.auth_headers() == {"Authorization": f"Bearer {provider.token}"} + assert subject.request("GET", https_server.url + "/inventory").json() == {"items": [123]} + assert subject.auth_headers() == {"Authorization": f"Bearer {provider.token}"} + + assert provider.accepted == 1 + exchange, resource = https_server.requests + assert exchange[1] == provider.profile["path"] + assert resource[2]["Authorization"] == f"Bearer {provider.token}" + assert provider.secret not in str(resource) + assert "Basic " not in str(resource) + + +def test_provider_contract_renews_with_rotated_credentials(provider, https_server, monkeypatch): + current_secret = [provider.secret] + loaded = [] + original_monotonic = time.monotonic + elapsed = [0] + monkeypatch.setattr(time, "monotonic", lambda: original_monotonic() + elapsed[0]) + + def load_secret(): + loaded.append(current_secret[0]) + return current_secret[0] + + subject = provider.client(client_secret=load_secret) + first_headers = subject.auth_headers() + lifetime = provider.profile["response"]["expires_in"] + elapsed[0] = lifetime - 31 + assert subject.auth_headers() == first_headers + assert len(loaded) == 1 + + provider.secret += "-rotated" + current_secret[0] = provider.secret + provider.generation += 1 + elapsed[0] = lifetime - 29 + + assert subject.auth_headers() == {"Authorization": f"Bearer {provider.token}"} + assert subject.auth_headers() != first_headers + assert len(loaded) == 2 + assert loaded[0] != loaded[1] + assert provider.accepted == 2 + assert len(https_server.requests) == 2 + + +def test_provider_contract_rejects_credentials_without_retry_or_leaking_details(provider, https_server, caplog): + subject = provider.client(client_secret="incorrect-test-secret") + + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + + assert not error.value.retryable + assert error.value.__context__ is None + assert error.value.__cause__ is None + assert provider.accepted == 0 + assert len(https_server.requests) == 1 + diagnostic = "".join(traceback.format_exception(error.type, error.value, error.tb)) + caplog.text + for private in (provider.secret, "incorrect-test-secret", provider.profile["error"]["error_description"]): + assert private not in diagnostic + + +def test_provider_contract_rejects_wrong_resource_or_scope(provider, https_server): + # Keycloak has no explicit scope here; an unassigned scope is rejected by this fixture. + subject = provider.client(scopes=["wrong-scope"], audience=None) + + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + + assert not error.value.retryable + assert provider.accepted == 0 + assert len(https_server.requests) == 1 + + +def test_provider_contract_rejects_wrong_authentication_method_without_fallback(provider, https_server): + other_method = "client_secret_post" if provider.auth_method == "client_secret_basic" else "client_secret_basic" + + with pytest.raises(TokenExchangeError) as error: + provider.client(auth_method=other_method).auth_headers() + + assert not error.value.retryable + assert provider.accepted == 0 + assert len(https_server.requests) == 1 + + +@pytest.mark.parametrize("post_providers", [(), tuple(PROFILES), ("entra", "keycloak")], ids=["basic", "post", "mixed"]) +def test_provider_contracts_keep_concurrent_clients_and_failures_isolated(https_server, post_providers): + providers = [ + Provider( + https_server, + name, + auth_method="client_secret_post" if name in post_providers else "client_secret_basic", + ) + for name in PROFILES + ] + clients = [provider.client() for provider in providers] + barrier = threading.Barrier(8) + + def acquire(index): + barrier.wait(timeout=5) + return index, clients[index].auth_headers() + + with ThreadPoolExecutor(max_workers=8) as executor: + results = list(executor.map(acquire, [0, 1, 2, 3] * 2)) + + for index, headers in results: + assert headers == {"Authorization": f"Bearer {providers[index].token}"} + assert all(provider.accepted == 1 for provider in providers) + assert Counter(request[1] for request in https_server.requests) == { + provider.profile["path"]: 1 for provider in providers + } + + providers[0].forced_error = 403, {"error": "access_denied"} + with pytest.raises(TokenExchangeError): + providers[0].client().auth_headers() + for index in range(1, 4): + assert clients[index].auth_headers() == {"Authorization": f"Bearer {providers[index].token}"} + assert len(https_server.requests) == 5 + + +def test_okta_management_scope_requires_an_unsupported_authentication_method(https_server): + # The Okta org authorization server requires private_key_jwt for these scopes. + # A synthetic rejection ensures the client does not retry or switch auth methods. + path = "/oauth2/v1/token" + https_server.serve( + path, + {"error": "invalid_client", "error_description": "private_key_jwt is required for this client"}, + status=401, + ) + subject = OAuth2Client( + token_url=https_server.url + path, + client_id="service-client", + client_secret="test-secret", + scopes=["okta.users.read"], + ) + + with pytest.raises(TokenExchangeError) as error: + subject.auth_headers() + + assert not error.value.retryable + assert len(https_server.requests) == 1 + request = https_server.requests[0] + assert request[2]["Authorization"].startswith("Basic ") + assert parse_qs(request[3].decode()) == { + "grant_type": ["client_credentials"], + "scope": ["okta.users.read"], + } diff --git a/uv.lock b/uv.lock index c500d2c1a38..2975fbd86f8 100644 --- a/uv.lock +++ b/uv.lock @@ -333,6 +333,9 @@ kafka-consumer-avro = [ kafka-consumer-protobuf = [ { name = "protobuf" }, ] +oauth2 = [ + { name = "urllib3" }, +] parser = [ { name = "pydantic" }, ] @@ -431,9 +434,10 @@ requires-dist = [ { name = "typing-extensions", specifier = ">=4.11.0,<5.0.0" }, { name = "urllib3", marker = "extra == 'all'", specifier = ">=2.8.0,<3.0.0" }, { name = "urllib3", marker = "extra == 'jwt'", specifier = ">=2.8.0,<3.0.0" }, + { name = "urllib3", marker = "extra == 'oauth2'", specifier = ">=2.8.0,<3.0.0" }, { name = "valkey-glide", marker = "extra == 'valkey'", specifier = ">=1.3.5,<3.0" }, ] -provides-extras = ["parser", "validation", "tracer", "redis", "valkey", "jwt", "all", "aws-sdk", "datadog", "datamasking", "kafka-consumer-avro", "kafka-consumer-protobuf"] +provides-extras = ["parser", "validation", "tracer", "redis", "valkey", "jwt", "oauth2", "all", "aws-sdk", "datadog", "datamasking", "kafka-consumer-avro", "kafka-consumer-protobuf"] [package.metadata.requires-dev] dev = [