diff --git a/aws_lambda_powertools/event_handler/api_gateway.py b/aws_lambda_powertools/event_handler/api_gateway.py index a323cf67d56..a764f0ca3a6 100644 --- a/aws_lambda_powertools/event_handler/api_gateway.py +++ b/aws_lambda_powertools/event_handler/api_gateway.py @@ -3035,21 +3035,18 @@ def include_router(self, router: Router, prefix: str | None = None) -> None: An optional prefix to be added to the originally defined rule """ - # Add reference to parent ApiGatewayResolver to support use cases where people subclass it to add custom logic - router.api_resolver = self - logger.debug("Merging App context with Router context") self.context.update(**router.context) + # Delegate request state to the resolver after preserving the router context. + router.api_resolver = self + logger.debug("Appending Router middlewares into App middlewares.") self._router_middlewares = self._router_middlewares + router._router_middlewares logger.debug("Appending Router exception_handler into App exception_handler.") self.exception_handler_manager.update_exception_handlers(router._exception_handlers) - # use pointer to allow context clearance after event is processed e.g., resolve(evt, ctx) - router.context = self.context - # Iterate through the routes defined in the router to configure and apply middlewares for each route for route, func in router._routes.items(): new_route = route @@ -3112,9 +3109,48 @@ def __init__(self): self._routes: dict[tuple, Callable] = {} self._routes_with_middleware: dict[tuple, list[Callable]] = {} self.api_resolver: BaseRouter | None = None - self.context = {} # early init as customers might add context before event resolution + self._context: dict = {} # early init as customers might add context before event resolution self._exception_handlers: dict[type, Callable] = {} + @property + def current_event(self) -> BaseProxyEvent: + if self.api_resolver is not None: + return self.api_resolver.current_event + return BaseRouter.current_event + + @current_event.setter + def current_event(self, value: BaseProxyEvent) -> None: + if self.api_resolver is not None: + self.api_resolver.current_event = value + else: + BaseRouter.current_event = value + + @property + def lambda_context(self) -> LambdaContext: + if self.api_resolver is not None: + return self.api_resolver.lambda_context + return BaseRouter.lambda_context + + @lambda_context.setter + def lambda_context(self, value: LambdaContext) -> None: + if self.api_resolver is not None: + self.api_resolver.lambda_context = value + else: + BaseRouter.lambda_context = value + + @property + def context(self) -> dict: + if self.api_resolver is not None: + return self.api_resolver.context + return self._context + + @context.setter + def context(self, value: dict) -> None: + if self.api_resolver is not None: + self.api_resolver.context = value + else: + self._context = value + def route( self, rule: str, diff --git a/aws_lambda_powertools/event_handler/http_resolver.py b/aws_lambda_powertools/event_handler/http_resolver.py index 7206dbe6f1c..dab644431b1 100644 --- a/aws_lambda_powertools/event_handler/http_resolver.py +++ b/aws_lambda_powertools/event_handler/http_resolver.py @@ -1,6 +1,8 @@ from __future__ import annotations import base64 +from contextvars import ContextVar +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Callable from urllib.parse import parse_qs @@ -13,6 +15,8 @@ from aws_lambda_powertools.utilities.data_classes.common import BaseProxyEvent if TYPE_CHECKING: + from collections.abc import Mapping, MutableMapping + from aws_lambda_powertools.shared.cookies import Cookie @@ -95,7 +99,7 @@ def _from_dict(cls, data: dict[str, Any]) -> HttpProxyEvent: return instance @classmethod - def from_asgi(cls, scope: dict[str, Any], body: bytes | None = None) -> HttpProxyEvent: + def from_asgi(cls, scope: Mapping[str, Any], body: bytes | None = None) -> HttpProxyEvent: """ Create an HttpProxyEvent from an ASGI scope dict. @@ -159,6 +163,14 @@ def get_remaining_time_in_millis(self) -> int: # pragma: no cover return 300000 # 5 minutes +@dataclass +class _RequestState: + event: BaseProxyEvent | None = None + lambda_context: Any = None + context: dict = field(default_factory=dict) + processed_stack_frames: list[str] = field(default_factory=list) + + class HttpResolverLocal(ApiGatewayResolver): """ ASGI-compatible HTTP resolver. @@ -204,6 +216,8 @@ def __init__( strip_prefixes: list[str | Any] | None = None, enable_validation: bool = False, ): + self._startup_state = _RequestState() + self._request_state: ContextVar[_RequestState | None] = ContextVar("local_http_request", default=None) super().__init__( proxy_type=ProxyEventType.APIGatewayProxyEvent, # Use REST API format internally cors=cors, @@ -212,7 +226,46 @@ def __init__( strip_prefixes=strip_prefixes, enable_validation=enable_validation, ) - self._is_async_mode = False + + @property + def _state(self) -> _RequestState: + return self._request_state.get() or self._startup_state + + # Powertools declares these as mutable attributes. Properties preserve that + # interface while directing each task to its own state. asyncio.to_thread + # propagates the ContextVar, so middleware sees the same request dictionary. + @property + def current_event(self) -> BaseProxyEvent: + # Preserve the inherited synchronous resolve() path outside ASGI calls. + return self._state.event or BaseRouter.current_event + + @current_event.setter + def current_event(self, value: BaseProxyEvent) -> None: + self._state.event = value + + @property + def lambda_context(self) -> Any: + return self._state.lambda_context or BaseRouter.lambda_context + + @lambda_context.setter + def lambda_context(self, value: Any) -> None: + self._state.lambda_context = value + + @property + def context(self) -> dict: + return self._state.context + + @context.setter + def context(self, value: dict) -> None: + self._state.context = value + + @property + def processed_stack_frames(self) -> list[str]: + return self._state.processed_stack_frames + + @processed_stack_frames.setter + def processed_stack_frames(self, value: list[str]) -> None: + self._state.processed_stack_frames = value def _to_proxy_event(self, event: dict) -> BaseProxyEvent: """Convert event dict to HttpProxyEvent.""" @@ -234,7 +287,7 @@ async def _resolve_async(self) -> dict: # type: ignore[override] response_builder = await super()._resolve_async() return response_builder.build(self.current_event, self._cors) - async def asgi_handler(self, scope: dict, receive: Callable, send: Callable) -> None: + async def asgi_handler(self, scope: MutableMapping[str, Any], receive: Callable, send: Callable) -> None: """ ASGI interface - allows running with uvicorn/hypercorn/etc. @@ -274,25 +327,27 @@ async def asgi_handler(self, scope: dict, receive: Callable, send: Callable) -> # Create mock Lambda context context: Any = MockLambdaContext() - # Set up resolver state (similar to resolve()) - BaseRouter.current_event = self._to_proxy_event(event._data) - BaseRouter.lambda_context = context - - self._is_async_mode = True - + # Never write BaseRouter's class attributes: another ASGI request may + # enter while validation or the handler is awaiting I/O. + state = _RequestState( + event=self._to_proxy_event(event._data), + lambda_context=context, + context=self._startup_state.context.copy(), + ) + token = self._request_state.set(state) try: - # Use async resolve response = await self._resolve_async() finally: - self._is_async_mode = False - self.clear_context() + # Reset only this task's binding. Middleware threads may still be + # unwinding after cancellation and retain their request's state. + self._request_state.reset(token) # Send HTTP response await self._send_response(send, response) async def __call__( # type: ignore[override] self, - scope: dict, + scope: MutableMapping[str, Any], receive: Callable, send: Callable, ) -> None: diff --git a/tests/functional/event_handler/_pydantic/test_http_resolver_pydantic.py b/tests/functional/event_handler/_pydantic/test_http_resolver_pydantic.py index e088527e359..9f66784d579 100644 --- a/tests/functional/event_handler/_pydantic/test_http_resolver_pydantic.py +++ b/tests/functional/event_handler/_pydantic/test_http_resolver_pydantic.py @@ -10,6 +10,7 @@ from pydantic import BaseModel, Field from aws_lambda_powertools.event_handler import HttpResolverLocal +from aws_lambda_powertools.event_handler.api_gateway import BaseRouter, Router from aws_lambda_powertools.event_handler.http_resolver import MockLambdaContext from aws_lambda_powertools.event_handler.openapi.params import Query @@ -365,3 +366,177 @@ def create_user(user: UserModel) -> UserResponse: # THEN schema includes 422 response post_operation = schema.paths["/users"].post assert 422 in post_operation.responses + + +async def _post_concurrent_request(app, name: str | None): + payload = {} if name is None else {"name": name, "age": 30} + scope = { + "type": "http", + "method": "POST", + "path": "/concurrent", + "headers": [(b"content-type", b"application/json"), (b"x-request-name", (name or "invalid").encode())], + "query_string": b"", + } + send, captured = make_asgi_send() + await asyncio.wait_for(app(scope, make_asgi_receive(json.dumps(payload).encode()), send), timeout=5) + return captured["status_code"], json.loads(captured["body"]) + + +def test_concurrent_asgi_validation_preserves_each_request_body(): + # GIVEN one local application using native request validation + app = HttpResolverLocal(enable_validation=True) + + @app.post("/concurrent") + async def echo(user: UserModel) -> dict: + await asyncio.sleep(0) + return {"name": user.name} + + async def scenario(): + # WHEN distinct bodies are submitted concurrently + return await asyncio.gather(*(_post_concurrent_request(app, str(i)) for i in range(6))) + + # THEN each caller receives its own input + assert asyncio.run(scenario()) == [(200, {"name": str(i)}) for i in range(6)] + + +def _make_included_router_app(): + app = HttpResolverLocal(enable_validation=True) + router = Router() + + @router.post("/concurrent") + async def echo(user: UserModel) -> dict: + app.append_context(name=user.name) + await asyncio.sleep(0) + return { + "name": router.context["name"], + "header": router.current_event.headers["x-request-name"], + "request_id": router.lambda_context.aws_request_id, + } + + app.include_router(router) + return app + + +def test_asgi_included_router_uses_request_state(): + app = _make_included_router_app() + + assert asyncio.run(_post_concurrent_request(app, "one")) == ( + 200, + {"name": "one", "header": "one", "request_id": "local-request-id"}, + ) + + +def test_concurrent_asgi_included_router_preserves_request_state(): + app = _make_included_router_app() + + async def scenario(): + return await asyncio.gather(*(_post_concurrent_request(app, str(i)) for i in range(6))) + + assert asyncio.run(scenario()) == [ + (200, {"name": str(i), "header": str(i), "request_id": "local-request-id"}) for i in range(6) + ] + + +def test_router_state_before_inclusion(monkeypatch): + app = HttpResolverLocal() + router = Router() + event = app._to_proxy_event( + { + "httpMethod": "GET", + "path": "/router", + "headers": {}, + "queryStringParameters": {}, + "multiValueQueryStringParameters": {}, + "body": None, + }, + ) + context = MockLambdaContext() + + monkeypatch.setattr(BaseRouter, "current_event", None, raising=False) + monkeypatch.setattr(BaseRouter, "lambda_context", None, raising=False) + + router.current_event = event + router.lambda_context = context + router.context = {"source": "router"} + + assert router.current_event is event + assert router.lambda_context is context + assert router.context == {"source": "router"} + + +def test_included_router_delegates_state_writes(): + app = HttpResolverLocal() + router = Router() + router.context = {"before": "include"} + app.include_router(router) + + event = app._to_proxy_event( + { + "httpMethod": "GET", + "path": "/router", + "headers": {}, + "queryStringParameters": {}, + "multiValueQueryStringParameters": {}, + "body": None, + }, + ) + context = MockLambdaContext() + + router.current_event = event + router.lambda_context = context + router.context = {"source": "resolver"} + + assert app.current_event is event + assert app.lambda_context is context + assert app.context == {"source": "resolver"} + + +@pytest.mark.parametrize("interruption", ["invalid", "cancelled"]) +def test_interrupted_asgi_request_does_not_clear_another_requests_state(interruption): + async def scenario(): + # GIVEN an active request that reads its context after awaiting I/O + app = HttpResolverLocal(enable_validation=True) + entered = {name: asyncio.Event() for name in ("first", "second", "later")} + release = {name: asyncio.Event() for name in entered} + pending = [] + + @app.post("/concurrent") + async def echo(user: UserModel) -> dict: + app.append_context(name=user.name) + entered[user.name].set() + await release[user.name].wait() + return {"name": app.context["name"], "header": app.current_event.headers["x-request-name"]} + + async def start(name): + task = asyncio.create_task(_post_concurrent_request(app, name)) + pending.append(task) + await asyncio.wait_for(entered[name].wait(), timeout=5) + return task + + try: + first = await start("first") + # WHEN another request fails validation or an overlapping request is cancelled + if interruption == "invalid": + status, _ = await _post_concurrent_request(app, None) + assert status == 422 + survivor, name = first, "first" + else: + survivor, name = await start("second"), "second" + first.cancel() + with pytest.raises(asyncio.CancelledError): + await first + + # THEN the surviving request retains both context and headers + release[name].set() + assert await survivor == (200, {"name": name, "header": name}) + release["later"].set() + assert await _post_concurrent_request(app, "later") == (200, {"name": "later", "header": "later"}) + finally: + for event in release.values(): + event.set() + for task in pending: + if not task.done(): + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + + asyncio.run(scenario())