Skip to content

Commit 5ca2e01

Browse files
committed
feat(auth): add JWT verification and OAuth client credentials
Add JWT verification, coordinated JWKS caching, API Gateway authorization, and OAuth client credentials with optional dependencies, documentation, examples, and tests. Include exception-safe claims cleanup, sanitized provider errors, and lazy imports for OAuth-only clients and static-key verification.
1 parent f872ab8 commit 5ca2e01

42 files changed

Lines changed: 4141 additions & 72 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.
Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,25 @@
1+
"""JWT verification and OAuth2 client credentials for AWS Lambda."""
2+
3+
from __future__ import annotations
4+
5+
import importlib
6+
from typing import TYPE_CHECKING
7+
8+
if TYPE_CHECKING:
9+
from aws_lambda_powertools.utilities.auth.oauth2 import OAuth2Client as OAuth2Client
10+
from aws_lambda_powertools.utilities.auth.verifier import JWTVerifier as JWTVerifier
11+
12+
__all__ = ["JWTVerifier", "OAuth2Client"]
13+
14+
15+
def __getattr__(name: str) -> object:
16+
modules = {"JWTVerifier": "verifier", "OAuth2Client": "oauth2"}
17+
if name in modules:
18+
value = getattr(importlib.import_module(f"{__name__}.{modules[name]}"), name)
19+
globals()[name] = value
20+
return value
21+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
22+
23+
24+
def __dir__() -> list[str]:
25+
return sorted(set(globals()) | set(__all__))
Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,79 @@
1+
from __future__ import annotations
2+
3+
from collections.abc import Mapping
4+
from typing import Any
5+
6+
from aws_lambda_powertools.utilities.auth._validation import string_list
7+
from aws_lambda_powertools.utilities.auth.exceptions import AuthError, InvalidClaimsError, InvalidTokenError
8+
9+
10+
class MissingTokenError(InvalidTokenError):
11+
"""No authorization header was supplied."""
12+
13+
14+
class ForbiddenError(AuthError):
15+
"""A verified caller does not have permission for this operation."""
16+
17+
18+
class InsufficientScopeError(ForbiddenError):
19+
"""A verified caller is missing a required scope."""
20+
21+
22+
def bearer_token(value: Any) -> str:
23+
if value is None:
24+
raise MissingTokenError()
25+
if not isinstance(value, str):
26+
raise InvalidTokenError()
27+
parts = value.split()
28+
if len(parts) != 2 or parts[0].lower() != "bearer":
29+
raise InvalidTokenError()
30+
return parts[1]
31+
32+
33+
def header_token(headers: Any, multi_value_headers: Any = None) -> str:
34+
values = _authorization_values(headers)
35+
multi_values = _authorization_values(multi_value_headers)
36+
if multi_values:
37+
entries = multi_values[0]
38+
if not isinstance(entries, list) or len(entries) != 1:
39+
raise InvalidTokenError()
40+
if values and values[0] != entries[0]:
41+
raise InvalidTokenError()
42+
return bearer_token(entries[0])
43+
return bearer_token(values[0] if values else None)
44+
45+
46+
def _authorization_values(headers: Any) -> list[Any]:
47+
if headers is None:
48+
return []
49+
if not isinstance(headers, Mapping):
50+
raise InvalidTokenError()
51+
values = [value for name, value in headers.items() if isinstance(name, str) and name.lower() == "authorization"]
52+
if len(values) > 1:
53+
raise InvalidTokenError()
54+
return values
55+
56+
57+
def valid_scope(value: str) -> bool:
58+
return bool(value) and all(33 <= ord(character) <= 126 and character not in {'"', "\\"} for character in value)
59+
60+
61+
def required_scopes(scopes: list[str] | None) -> tuple[str, ...]:
62+
values = string_list(scopes if scopes is not None else [])
63+
if not all(valid_scope(value) for value in values):
64+
raise ValueError("Scopes must be valid OAuth scope tokens")
65+
return values
66+
67+
68+
def enforce_scopes(claims: dict[str, Any], expected: tuple[str, ...]) -> None:
69+
value: Any = next((claims[name] for name in ("scope", "scp", "scopes") if name in claims), [])
70+
if isinstance(value, str):
71+
values = [part for part in value.split(" ") if part]
72+
elif isinstance(value, list):
73+
values = value
74+
else:
75+
raise InvalidClaimsError()
76+
if any(not isinstance(scope, str) or not valid_scope(scope) for scope in values):
77+
raise InvalidClaimsError()
78+
if not set(expected).issubset(values):
79+
raise InsufficientScopeError()
Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,122 @@
1+
from __future__ import annotations
2+
3+
import math
4+
import re
5+
from typing import TYPE_CHECKING, Any, Literal
6+
7+
from aws_lambda_powertools.utilities.auth._authorization import (
8+
ForbiddenError,
9+
bearer_token,
10+
enforce_scopes,
11+
header_token,
12+
required_scopes,
13+
)
14+
from aws_lambda_powertools.utilities.auth._validation import string_list
15+
from aws_lambda_powertools.utilities.auth.exceptions import InvalidClaimsError, InvalidTokenError
16+
from aws_lambda_powertools.utilities.data_classes.api_gateway_authorizer_event import APIGatewayAuthorizerResponseV2
17+
from aws_lambda_powertools.utilities.data_classes.common import DictWrapper
18+
19+
if TYPE_CHECKING:
20+
from aws_lambda_powertools.utilities.auth._base import Verifier
21+
22+
_ARN = re.compile(r"arn:[a-z0-9-]+:execute-api:[a-z0-9-]+:\d{12}:[a-z0-9]+/[^/]+/[A-Z]+/.*")
23+
24+
25+
def authorize_event(
26+
verifier: Verifier,
27+
event: dict[str, Any] | DictWrapper,
28+
scopes: list[str] | None,
29+
response_format: Literal["iam", "simple"],
30+
context_claims: list[str] | None,
31+
) -> dict[str, Any]:
32+
raw = event.raw_event if isinstance(event, DictWrapper) else event
33+
_validate_event(raw, response_format)
34+
arn = _request_arn(raw) if response_format == "iam" else None
35+
expected = required_scopes(scopes)
36+
selected = string_list(context_claims if context_claims is not None else [])
37+
if "claims" in selected:
38+
raise ValueError("claims is reserved in API Gateway authorizer context")
39+
claims = _verified_claims(verifier, raw, expected, require_principal=response_format == "iam")
40+
context = _context(claims, selected) if claims is not None else {}
41+
if response_format == "simple":
42+
return APIGatewayAuthorizerResponseV2(authorize=claims is not None, context=context).asdict()
43+
return _iam_response(claims, arn, context)
44+
45+
46+
def _validate_event(raw: dict[str, Any], response_format: str) -> None:
47+
if not isinstance(raw, dict) or raw.get("type") not in ("TOKEN", "REQUEST"):
48+
raise ValueError("An API Gateway TOKEN or REQUEST authorizer event is required")
49+
if response_format not in ("iam", "simple"):
50+
raise ValueError("response_format must be iam or simple")
51+
if response_format == "simple" and (raw.get("version") != "2.0" or raw["type"] != "REQUEST"):
52+
raise ValueError("Simple authorizer responses require HTTP API payload version 2.0")
53+
54+
55+
def _verified_claims(
56+
verifier: Verifier,
57+
raw: dict[str, Any],
58+
expected: tuple[str, ...],
59+
*,
60+
require_principal: bool,
61+
) -> dict[str, Any] | None:
62+
try:
63+
candidate = verifier.verify(_token(raw))
64+
enforce_scopes(candidate, expected)
65+
if require_principal:
66+
_validate_principal(candidate)
67+
return candidate
68+
except (InvalidTokenError, ForbiddenError):
69+
return None
70+
71+
72+
def _token(raw: dict[str, Any]) -> str:
73+
if raw["type"] == "TOKEN":
74+
return bearer_token(raw.get("authorizationToken"))
75+
return header_token(raw.get("headers"), raw.get("multiValueHeaders"))
76+
77+
78+
def _validate_principal(claims: dict[str, Any]) -> None:
79+
if not isinstance(claims.get("sub"), str) or not claims["sub"].strip():
80+
raise InvalidClaimsError()
81+
82+
83+
def _iam_response(claims: dict[str, Any] | None, arn: str | None, context: dict[str, Any]) -> dict[str, Any]:
84+
# Preserve the exact supplied resource, including its partition and encoded
85+
# path. Route builders normalize paths and cannot represent every ARN here.
86+
result: dict[str, Any] = {
87+
"principalId": claims["sub"] if claims is not None else "unauthorized",
88+
"policyDocument": {
89+
"Version": "2012-10-17",
90+
"Statement": [
91+
{
92+
"Action": "execute-api:Invoke",
93+
"Effect": "Allow" if claims is not None else "Deny",
94+
"Resource": [arn],
95+
},
96+
],
97+
},
98+
}
99+
if context:
100+
result["context"] = context
101+
return result
102+
103+
104+
def _request_arn(event: dict[str, Any]) -> str:
105+
arn = event.get("routeArn") if event.get("version") == "2.0" else event.get("methodArn")
106+
if (
107+
not isinstance(arn, str)
108+
or len(arn) > 512
109+
or not _ARN.fullmatch(arn)
110+
or any(character in arn for character in ("*", "?", "\r", "\n"))
111+
):
112+
raise ValueError("A concrete API Gateway method or route ARN of at most 512 characters is required")
113+
return arn
114+
115+
116+
def _context(claims: dict[str, Any], selected: tuple[str, ...]) -> dict[str, Any]:
117+
context = {}
118+
for name in selected:
119+
value = claims.get(name)
120+
if isinstance(value, (str, bool, int)) or isinstance(value, float) and math.isfinite(value):
121+
context[name] = value
122+
return context
Lines changed: 98 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,98 @@
1+
from __future__ import annotations
2+
3+
from abc import ABC, abstractmethod
4+
from typing import TYPE_CHECKING, Any, Literal
5+
6+
if TYPE_CHECKING:
7+
from collections.abc import Callable
8+
9+
from aws_lambda_powertools.event_handler import Response
10+
from aws_lambda_powertools.utilities.auth._middleware import AuthErrorContext, AuthMiddleware
11+
from aws_lambda_powertools.utilities.data_classes.common import DictWrapper
12+
13+
14+
class Verifier(ABC):
15+
"""Shared verification interface used by issuer-specific and routed verifiers."""
16+
17+
@abstractmethod
18+
def verify(self, token: str) -> dict[str, Any]:
19+
"""Return verified claims or raise an Auth utility error."""
20+
21+
@abstractmethod
22+
def prefetch(self) -> None:
23+
"""Populate remote key caches without accepting a token."""
24+
25+
def require(
26+
self,
27+
*,
28+
scopes: list[str] | None = None,
29+
authorize: Callable[[dict[str, Any]], bool] | None = None,
30+
on_error: Callable[[AuthErrorContext], Response] | None = None,
31+
) -> AuthMiddleware:
32+
"""Create Event Handler middleware enforcing token validity and all scopes.
33+
34+
Successful verification stores claims in ``app.context["claims"]``
35+
while the downstream middleware and handler execute. Claims are
36+
removed when they return or raise.
37+
Missing/invalid tokens return 401, missing permissions return 403, and
38+
unavailable signing keys return 503. A custom error callback replaces
39+
the response, never execution of the protected handler.
40+
41+
Parameters
42+
----------
43+
scopes : list[str], optional
44+
Every listed scope must be present in the token.
45+
authorize : Callable, optional
46+
Additional policy receiving verified claims; must return True.
47+
on_error : Callable, optional
48+
Receives status_code and headers and returns an Event Handler Response.
49+
50+
Examples
51+
--------
52+
```python
53+
@app.get("/orders", middlewares=[verifier.require(scopes=["orders:read"])])
54+
def orders():
55+
return {"subject": app.context["claims"]["sub"]}
56+
```
57+
"""
58+
from aws_lambda_powertools.utilities.auth._middleware import AuthMiddleware
59+
60+
return AuthMiddleware(self, scopes, authorize, on_error)
61+
62+
def authorize(
63+
self,
64+
event: dict[str, Any] | DictWrapper,
65+
*,
66+
scopes: list[str] | None = None,
67+
response_format: Literal["iam", "simple"] = "iam",
68+
context_claims: list[str] | None = None,
69+
) -> dict[str, Any]:
70+
"""Return an API Gateway authorizer response for the current request.
71+
72+
IAM allows require a nonempty ``sub`` and target the supplied ARN only.
73+
Simple responses require payload version 2.0 and must also be enabled
74+
in the Gateway deployment. Disable Gateway result caching when each
75+
request must be verified; this method cannot change Gateway's TTL.
76+
77+
Parameters
78+
----------
79+
event : dict | DictWrapper
80+
REST TOKEN/REQUEST or HTTP REQUEST authorizer event.
81+
scopes : list[str], optional
82+
Every listed scope must be present in the token.
83+
response_format : Literal["iam", "simple"]
84+
Response format configured in Gateway, by default iam.
85+
context_claims : list[str], optional
86+
Selected scalar claims to include; no claims are copied by default.
87+
88+
Examples
89+
--------
90+
```python
91+
return verifier.authorize(
92+
event, scopes=["orders:read"], response_format="iam", context_claims=["sub"],
93+
)
94+
```
95+
"""
96+
from aws_lambda_powertools.utilities.auth._authorizer import authorize_event
97+
98+
return authorize_event(self, event, scopes, response_format, context_claims)
Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,26 @@
1+
from __future__ import annotations
2+
3+
import time
4+
5+
from aws_lambda_powertools.utilities.auth._validation import finite_seconds
6+
7+
8+
class RequestError(Exception):
9+
"""Internal, credential-free transport failure."""
10+
11+
def __init__(self, *, retryable: bool = False) -> None:
12+
self.retryable = retryable
13+
super().__init__("Authentication endpoint request failed")
14+
15+
16+
class Deadline:
17+
"""One monotonic budget shared across a fetch and any subsequent requests."""
18+
19+
def __init__(self, seconds: float) -> None:
20+
self._expires_at = time.monotonic() + finite_seconds(seconds, positive=True)
21+
22+
def remaining(self) -> float:
23+
remaining = self._expires_at - time.monotonic()
24+
if remaining <= 0:
25+
raise RequestError(retryable=True)
26+
return remaining
Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
from __future__ import annotations
2+
3+
from functools import wraps
4+
from typing import TYPE_CHECKING, ParamSpec, TypeVar
5+
6+
from aws_lambda_powertools.utilities.auth.exceptions import AuthError
7+
8+
if TYPE_CHECKING:
9+
from collections.abc import Callable
10+
11+
_P = ParamSpec("_P")
12+
_T = TypeVar("_T")
13+
14+
15+
def sanitize_errors(operation: Callable[_P, _T]) -> Callable[_P, _T]:
16+
"""Detach provider exceptions before an Auth error leaves a public operation."""
17+
18+
@wraps(operation)
19+
def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _T:
20+
try:
21+
return operation(*args, **kwargs)
22+
except AuthError as error:
23+
# `raise ... from None` only suppresses display of the context.
24+
# Clear both references and use a bare re-raise so Python does not
25+
# attach the active exception again.
26+
error.__context__ = None
27+
error.__cause__ = None
28+
raise
29+
30+
return wrapper

0 commit comments

Comments
 (0)