diff --git a/aws_lambda_powertools/event_handler/api_gateway.py b/aws_lambda_powertools/event_handler/api_gateway.py index 25b34613f44..4fb4a6ce258 100644 --- a/aws_lambda_powertools/event_handler/api_gateway.py +++ b/aws_lambda_powertools/event_handler/api_gateway.py @@ -2560,7 +2560,10 @@ def resolve(self, event: Mapping[str, Any], context: LambdaContext) -> dict[str, BaseRouter.current_event = self._to_proxy_event(cast(dict, event)) BaseRouter.lambda_context = context - response = self._build_response(self._resolve()) + try: + response = self._build_response(self._resolve()) + finally: + self.clear_context() # Debug print Processed Middlewares if self._debug: @@ -2569,8 +2572,6 @@ def resolve(self, event: Mapping[str, Any], context: LambdaContext) -> dict[str, print("\n".join(self.processed_stack_frames)) print("======================") - self.clear_context() - return response async def resolve_async(self, event: Mapping[str, Any], context: LambdaContext) -> dict[str, Any]: @@ -2621,7 +2622,10 @@ def lambda_handler(event, context): BaseRouter.current_event = self._to_proxy_event(cast(dict, event)) BaseRouter.lambda_context = context - response = self._build_response(await self._resolve_async()) + try: + response = self._build_response(await self._resolve_async()) + finally: + self.clear_context() if self._debug: print("\nProcessed Middlewares:") @@ -2629,8 +2633,6 @@ def lambda_handler(event, context): print("\n".join(self.processed_stack_frames)) print("======================") - self.clear_context() - return response def _build_response(self, response_builder: ResponseBuilder) -> dict[str, Any]: @@ -3401,20 +3403,16 @@ def _build_response(self, response_builder: ResponseBuilder) -> dict[str, Any]: try: self._validate_response_size(response) except ResponseSizeExceededError as exc: - try: - # Resolved responses retain their route, including not-found and preflight responses. - handled_response = self._call_exception_handler(exc, cast(Route, response_builder.route)) - if handled_response is None: - raise - - handled_response.response = cast(Response, self._to_response(handled_response.response)) - response = super()._build_response(handled_response) - # Validate once more without invoking an exception handler recursively. - self._validate_response_size(response) - except Exception: - self.clear_context() + # Resolved responses retain their route, including not-found and preflight responses. + handled_response = self._call_exception_handler(exc, cast(Route, response_builder.route)) + if handled_response is None: raise + handled_response.response = cast(Response, self._to_response(handled_response.response)) + response = super()._build_response(handled_response) + # Validate once more without invoking an exception handler recursively. + self._validate_response_size(response) + return response @staticmethod diff --git a/tests/functional/auth_alpha/jwt/integrations/test_middleware.py b/tests/functional/auth_alpha/jwt/integrations/test_middleware.py index 97639182528..bde43f6887b 100644 --- a/tests/functional/auth_alpha/jwt/integrations/test_middleware.py +++ b/tests/functional/auth_alpha/jwt/integrations/test_middleware.py @@ -84,8 +84,7 @@ def public(): event["headers"] = {"authorization": "Bearer " + issue_token()} with pytest.raises(RuntimeError, match="handler failed"): app.resolve(event, {}) - assert "claims" not in app.context - assert app.context["application_value"] == "preserved" + assert app.context == {} event["headers"] = {} public_event = copy.deepcopy(event) diff --git a/tests/functional/event_handler/required_dependencies/test_api_gateway.py b/tests/functional/event_handler/required_dependencies/test_api_gateway.py index dec2e6d3b2b..65d1a50082d 100644 --- a/tests/functional/event_handler/required_dependencies/test_api_gateway.py +++ b/tests/functional/event_handler/required_dependencies/test_api_gateway.py @@ -1892,6 +1892,23 @@ def my_path(): assert app.context == {} +def test_route_context_is_cleared_when_handler_raises(): + # GIVEN a route that raises an exception without a registered exception handler + app = APIGatewayRestResolver() + app.append_context(is_admin=True) + + @app.get("/my/path") + def my_path(): + raise ValueError("boom") + + # WHEN event resolution kicks in + with pytest.raises(ValueError, match="boom"): + app.resolve(LOAD_GW_EVENT, {}) + + # THEN context should be cleared even though the route raised + assert app.context == {} + + def test_router_has_access_to_app_context(json_dump): # GIVEN a Router with registered routes app = ApiGatewayResolver() diff --git a/tests/functional/event_handler/required_dependencies/test_request.py b/tests/functional/event_handler/required_dependencies/test_request.py index b00ae6659ba..e1147f3805b 100644 --- a/tests/functional/event_handler/required_dependencies/test_request.py +++ b/tests/functional/event_handler/required_dependencies/test_request.py @@ -321,6 +321,24 @@ def handler(counter_id: str, request: Request): assert call_count == 3 +def test_request_is_not_reused_after_unhandled_exception(): + """A handler raising must not leave its Request cached for the next invocation.""" + app = APIGatewayRestResolver() + received: list[Request] = [] + + @app.get("/counters/") + def handler(counter_id: str, request: Request): + received.append(request) + raise ValueError(counter_id) + + for i in range(2): + event = _make_rest_event(f"/counters/{i}", path_parameters={"counter_id": str(i)}) + with pytest.raises(ValueError): + app(event, {}) + + assert [req.path_parameters for req in received] == [{"counter_id": "0"}, {"counter_id": "1"}] + + # --------------------------------------------------------------------------- # RuntimeError when accessed outside of request resolution # --------------------------------------------------------------------------- diff --git a/tests/functional/event_handler/required_dependencies/test_resolve_async.py b/tests/functional/event_handler/required_dependencies/test_resolve_async.py index e9b12ce2a2d..d15895f6a6c 100644 --- a/tests/functional/event_handler/required_dependencies/test_resolve_async.py +++ b/tests/functional/event_handler/required_dependencies/test_resolve_async.py @@ -456,6 +456,22 @@ async def get_lambda(): # THEN the context is cleared after resolution assert app.context == {} + def test_resolve_async_clears_context_when_handler_raises(self, public_resolver_and_event): + # GIVEN an async handler that raises without a registered exception handler + app, event, path = public_resolver_and_event + + @app.get(path) + async def get_lambda(): + app.append_context(custom_key="value") + raise ValueError("boom") + + # WHEN calling resolve_async + with pytest.raises(ValueError, match="boom"): + asyncio.run(app.resolve_async(event, MockLambdaContext())) + + # THEN the context is still cleared + assert app.context == {} + def test_resolve_async_not_found(self, public_resolver_and_event): # GIVEN no matching route app, event, _path = public_resolver_and_event