From 6a18074092d772dd4e2dc7790f79687a14819036 Mon Sep 17 00:00:00 2001 From: Emanuele Giaquinta Date: Thu, 8 Oct 2026 19:21:38 +0300 Subject: [PATCH] refactor(typing): add type hints to middleware-based decorators lambda_handler_decorator returns an untyped Callable, so handlers decorated with event_parser, idempotent, kafka_consumer and validator were typed as Any. event_parser, kafka_consumer and validator are typed with overloads. idempotent, which can only be called with options, is instead cast to a Protocol, following the existing event_source pattern. Decorated handlers are typed with a shared LambdaHandler Protocol, so they can be called with keyword arguments. --- aws_lambda_powertools/shared/types.py | 12 +++- .../utilities/data_classes/event_source.py | 5 +- .../utilities/idempotency/idempotency.py | 32 +++++++++- .../utilities/kafka/kafka_consumer.py | 27 +++++++- .../utilities/parser/parser.py | 54 +++++++++++++++- .../utilities/validation/validator.py | 63 ++++++++++++++++++- .../test_disabling_idempotency_utility.py | 7 ++- ..._fields_contain_json_pydantic_validator.py | 6 +- 8 files changed, 193 insertions(+), 13 deletions(-) diff --git a/aws_lambda_powertools/shared/types.py b/aws_lambda_powertools/shared/types.py index aeafb378dab..701904f0c08 100644 --- a/aws_lambda_powertools/shared/types.py +++ b/aws_lambda_powertools/shared/types.py @@ -1,4 +1,14 @@ from collections.abc import Callable -from typing import Any, TypeVar +from typing import Any, Protocol, TypeVar AnyCallableT = TypeVar("AnyCallableT", bound=Callable[..., Any]) + +EventT_contra = TypeVar("EventT_contra", contravariant=True) +ContextT_contra = TypeVar("ContextT_contra", contravariant=True) +ReturnT_co = TypeVar("ReturnT_co", covariant=True) + + +class LambdaHandler(Protocol[EventT_contra, ContextT_contra, ReturnT_co]): + """Lambda handler, callable with `event` and `context` as positional or keyword arguments.""" + + def __call__(self, event: EventT_contra, context: ContextT_contra) -> ReturnT_co: ... diff --git a/aws_lambda_powertools/utilities/data_classes/event_source.py b/aws_lambda_powertools/utilities/data_classes/event_source.py index 58c8be968f5..9691a87065c 100644 --- a/aws_lambda_powertools/utilities/data_classes/event_source.py +++ b/aws_lambda_powertools/utilities/data_classes/event_source.py @@ -7,6 +7,7 @@ if TYPE_CHECKING: from collections.abc import Callable + from aws_lambda_powertools.shared.types import LambdaHandler from aws_lambda_powertools.utilities.data_classes.common import DictWrapper from aws_lambda_powertools.utilities.typing import LambdaContext @@ -15,7 +16,7 @@ class _EventSourceDecorator(Protocol): - """Annotation of event_source, that lambda_handler_decorator erases at runtime.""" + """Type of `event_source`, as `lambda_handler_decorator` returns an untyped `Callable`.""" def __call__( self, @@ -23,7 +24,7 @@ def __call__( data_class: type[DataClassT], ) -> Callable[ [Callable[[DataClassT, LambdaContext], OutputT]], - Callable[[dict[str, Any], LambdaContext], OutputT], + LambdaHandler[dict[str, Any], LambdaContext, OutputT], ]: ... diff --git a/aws_lambda_powertools/utilities/idempotency/idempotency.py b/aws_lambda_powertools/utilities/idempotency/idempotency.py index 7dd0f0fd198..62a8994889c 100644 --- a/aws_lambda_powertools/utilities/idempotency/idempotency.py +++ b/aws_lambda_powertools/utilities/idempotency/idempotency.py @@ -9,7 +9,7 @@ import os import warnings from inspect import isclass -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, Protocol, TypeVar, cast from aws_lambda_powertools.middleware_factory import lambda_handler_decorator from aws_lambda_powertools.shared import constants @@ -25,6 +25,7 @@ if TYPE_CHECKING: from collections.abc import Callable + from aws_lambda_powertools.shared.types import LambdaHandler from aws_lambda_powertools.utilities.idempotency.persistence.base import ( BasePersistenceLayer, ) @@ -34,9 +35,33 @@ logger = logging.getLogger(__name__) +EventT = TypeVar("EventT") +ContextT = TypeVar("ContextT", bound="LambdaContext | DurableContextProtocol") + + +class _IdempotentHandlerDecorator(Protocol): + # The return type is Any, as responses replayed from the persistence store are deserialized JSON + def __call__( + self, + handler: Callable[[EventT, ContextT], Any], + /, + ) -> LambdaHandler[EventT, ContextT, Any]: ... + + +class _IdempotentDecorator(Protocol): + """Type of `idempotent`, as `lambda_handler_decorator` returns an untyped `Callable`.""" + + def __call__( + self, + *, + persistence_store: BasePersistenceLayer, + config: IdempotencyConfig | None = None, + key_prefix: str | None = None, + ) -> _IdempotentHandlerDecorator: ... + @lambda_handler_decorator -def idempotent( +def _idempotent( handler: Callable[[Any, LambdaContext | DurableContextProtocol], Any], event: dict[str, Any], context: LambdaContext | DurableContextProtocol, @@ -117,6 +142,9 @@ def handler(event, context): return idempotency_handler.handle(is_replay=is_replay) +idempotent = cast(_IdempotentDecorator, _idempotent) + + def idempotent_function( function: AnyCallableT | None = None, *, diff --git a/aws_lambda_powertools/utilities/kafka/kafka_consumer.py b/aws_lambda_powertools/utilities/kafka/kafka_consumer.py index b4bacea545e..5dbca338a98 100644 --- a/aws_lambda_powertools/utilities/kafka/kafka_consumer.py +++ b/aws_lambda_powertools/utilities/kafka/kafka_consumer.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Protocol, TypeVar, overload from aws_lambda_powertools.middleware_factory import lambda_handler_decorator from aws_lambda_powertools.utilities.kafka.consumer_records import ConsumerRecords @@ -8,9 +8,34 @@ if TYPE_CHECKING: from collections.abc import Callable + from aws_lambda_powertools.shared.types import LambdaHandler from aws_lambda_powertools.utilities.kafka.schema_config import SchemaConfig from aws_lambda_powertools.utilities.typing import LambdaContext +ReturnT = TypeVar("ReturnT") + + +class _KafkaConsumerHandlerDecorator(Protocol): + def __call__( + self, + handler: Callable[[ConsumerRecords, LambdaContext], ReturnT], + /, + ) -> LambdaHandler[dict[str, Any], LambdaContext, ReturnT]: ... + + +@overload +def kafka_consumer( + handler: Callable[[ConsumerRecords, LambdaContext], ReturnT], + /, +) -> LambdaHandler[dict[str, Any], LambdaContext, ReturnT]: ... + + +@overload +def kafka_consumer( + *, + schema_config: SchemaConfig | None = None, +) -> _KafkaConsumerHandlerDecorator: ... + @lambda_handler_decorator def kafka_consumer( diff --git a/aws_lambda_powertools/utilities/parser/parser.py b/aws_lambda_powertools/utilities/parser/parser.py index 652cc1ebf1e..74e2d37b2f2 100644 --- a/aws_lambda_powertools/utilities/parser/parser.py +++ b/aws_lambda_powertools/utilities/parser/parser.py @@ -9,7 +9,7 @@ import logging import typing -from typing import TYPE_CHECKING, Any, Callable, overload +from typing import TYPE_CHECKING, Any, Callable, Protocol, TypeVar, overload from pydantic import PydanticSchemaGenerationError @@ -21,12 +21,62 @@ ) if TYPE_CHECKING: - from aws_lambda_powertools.utilities.parser.envelopes.base import Envelope + from pydantic import TypeAdapter + + from aws_lambda_powertools.shared.types import LambdaHandler + from aws_lambda_powertools.utilities.parser.envelopes.base import BaseEnvelope, Envelope from aws_lambda_powertools.utilities.parser.types import EventParserReturnType, T from aws_lambda_powertools.utilities.typing import LambdaContext logger = logging.getLogger(__name__) +ModelT_co = TypeVar("ModelT_co", covariant=True) +ModelT = TypeVar("ModelT") +ReturnT = TypeVar("ReturnT") + + +class _EventParserModelHandlerDecorator(Protocol[ModelT_co]): + def __call__( + self, + handler: Callable[[ModelT_co, LambdaContext], ReturnT], + /, + ) -> LambdaHandler[dict[str, Any] | str, LambdaContext, ReturnT]: ... + + +class _EventParserHandlerDecorator(Protocol): + # The handler's event type is Any, as it can't be derived from the arguments, e.g. + # - @event_parser() + # - @event_parser(model=Annotated[A | B, Field(discriminator="kind")]) + # - @event_parser(model=A, envelope=SqsEnvelope) + def __call__( + self, + handler: Callable[[Any, LambdaContext], ReturnT], + /, + ) -> LambdaHandler[dict[str, Any] | str, LambdaContext, ReturnT]: ... + + +@overload +def event_parser( + handler: Callable[[Any, LambdaContext], ReturnT], + /, +) -> LambdaHandler[dict[str, Any] | str, LambdaContext, ReturnT]: ... + + +@overload +def event_parser( + *, + model: type[ModelT] | TypeAdapter[ModelT], + envelope: None = None, +) -> _EventParserModelHandlerDecorator[ModelT]: ... + + +@overload +def event_parser( + *, + model: Any = None, + envelope: type[BaseEnvelope] | None = None, +) -> _EventParserHandlerDecorator: ... + @lambda_handler_decorator def event_parser( diff --git a/aws_lambda_powertools/utilities/validation/validator.py b/aws_lambda_powertools/utilities/validation/validator.py index bd9bb0db738..fc12224b3ce 100644 --- a/aws_lambda_powertools/utilities/validation/validator.py +++ b/aws_lambda_powertools/utilities/validation/validator.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Protocol, TypeVar, overload from aws_lambda_powertools.middleware_factory import lambda_handler_decorator from aws_lambda_powertools.utilities import jmespath_utils @@ -10,8 +10,69 @@ if TYPE_CHECKING: from collections.abc import Callable + from aws_lambda_powertools.shared.types import LambdaHandler + logger = logging.getLogger(__name__) +EventT = TypeVar("EventT") +ContextT = TypeVar("ContextT") +ReturnT = TypeVar("ReturnT") + + +class _ValidatorHandlerDecorator(Protocol): + def __call__( + self, + handler: Callable[[EventT, ContextT], ReturnT], + /, + ) -> LambdaHandler[EventT, ContextT, ReturnT]: ... + + +class _ValidatorEnvelopeHandlerDecorator(Protocol): + # The handler's event type is Any, as it can't be derived from the arguments, e.g. + # - @validator(envelope=envelopes.SQS) + def __call__( + self, + handler: Callable[[Any, ContextT], ReturnT], + /, + ) -> LambdaHandler[dict[str, Any] | str, ContextT, ReturnT]: ... + + +@overload +def validator( + handler: Callable[[EventT, ContextT], ReturnT], + /, +) -> LambdaHandler[EventT, ContextT, ReturnT]: ... + + +@overload +def validator( + *, + inbound_schema: dict | None = None, + inbound_formats: dict | None = None, + inbound_handlers: dict | None = None, + inbound_provider_options: dict | None = None, + outbound_schema: dict | None = None, + outbound_formats: dict | None = None, + outbound_handlers: dict | None = None, + outbound_provider_options: dict | None = None, +) -> _ValidatorHandlerDecorator: ... + + +@overload +def validator( + *, + inbound_schema: dict | None = None, + inbound_formats: dict | None = None, + inbound_handlers: dict | None = None, + inbound_provider_options: dict | None = None, + outbound_schema: dict | None = None, + outbound_formats: dict | None = None, + outbound_handlers: dict | None = None, + outbound_provider_options: dict | None = None, + envelope: str, + jmespath_options: dict | None = None, +) -> _ValidatorEnvelopeHandlerDecorator: ... + @lambda_handler_decorator def validator( diff --git a/examples/idempotency/tests/test_disabling_idempotency_utility.py b/examples/idempotency/tests/test_disabling_idempotency_utility.py index 3aba8a090c8..484f1e1fc0c 100644 --- a/examples/idempotency/tests/test_disabling_idempotency_utility.py +++ b/examples/idempotency/tests/test_disabling_idempotency_utility.py @@ -1,11 +1,14 @@ from dataclasses import dataclass +from typing import cast import app_test_disabling_idempotency_utility import pytest +from aws_lambda_powertools.utilities.typing import LambdaContext + @dataclass -class LambdaContext: +class FakeLambdaContext: function_name: str = "test" memory_limit_in_mb: int = 128 invoked_function_arn: str = "arn:aws:lambda:eu-west-1:809313241:function:test" @@ -17,7 +20,7 @@ def get_remaining_time_in_millis(self) -> int: @pytest.fixture def lambda_context() -> LambdaContext: - return LambdaContext() + return cast(LambdaContext, FakeLambdaContext()) def test_idempotent_lambda_handler(monkeypatch, lambda_context: LambdaContext): diff --git a/examples/parser/src/string_fields_contain_json_pydantic_validator.py b/examples/parser/src/string_fields_contain_json_pydantic_validator.py index 5c19606736d..f0054114e05 100644 --- a/examples/parser/src/string_fields_contain_json_pydantic_validator.py +++ b/examples/parser/src/string_fields_contain_json_pydantic_validator.py @@ -3,13 +3,14 @@ import json from typing import TYPE_CHECKING, Any +from pydantic import field_validator + from aws_lambda_powertools.utilities.parser import BaseEnvelope, BaseModel, event_parser from aws_lambda_powertools.utilities.parser.functions import ( _parse_and_validate_event, _retrieve_or_set_model_from_cache, ) from aws_lambda_powertools.utilities.typing import LambdaContext -from aws_lambda_powertools.utilities.validation import validator if TYPE_CHECKING: from aws_lambda_powertools.utilities.parser.types import T @@ -23,7 +24,8 @@ class CancelOrder(BaseModel): class CancelOrderModel(BaseModel): body: CancelOrder - @validator("body", pre=True) + @field_validator("body", mode="before") + @classmethod def transform_body_to_dict(cls, value): return json.loads(value) if isinstance(value, str) else value