Skip to content

Commit 7d3d400

Browse files
committed
fix(event_handler): isolate local ASGI request state
1 parent 2f69b41 commit 7d3d400

2 files changed

Lines changed: 150 additions & 13 deletions

File tree

‎aws_lambda_powertools/event_handler/http_resolver.py‎

Lines changed: 68 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
from __future__ import annotations
22

33
import base64
4+
from contextvars import ContextVar
5+
from dataclasses import dataclass, field
46
from typing import TYPE_CHECKING, Any, Callable
57
from urllib.parse import parse_qs
68

@@ -13,6 +15,8 @@
1315
from aws_lambda_powertools.utilities.data_classes.common import BaseProxyEvent
1416

1517
if TYPE_CHECKING:
18+
from collections.abc import Mapping, MutableMapping
19+
1620
from aws_lambda_powertools.shared.cookies import Cookie
1721

1822

@@ -95,7 +99,7 @@ def _from_dict(cls, data: dict[str, Any]) -> HttpProxyEvent:
9599
return instance
96100

97101
@classmethod
98-
def from_asgi(cls, scope: dict[str, Any], body: bytes | None = None) -> HttpProxyEvent:
102+
def from_asgi(cls, scope: Mapping[str, Any], body: bytes | None = None) -> HttpProxyEvent:
99103
"""
100104
Create an HttpProxyEvent from an ASGI scope dict.
101105
@@ -159,6 +163,14 @@ def get_remaining_time_in_millis(self) -> int: # pragma: no cover
159163
return 300000 # 5 minutes
160164

161165

166+
@dataclass
167+
class _RequestState:
168+
event: BaseProxyEvent | None = None
169+
lambda_context: Any = None
170+
context: dict = field(default_factory=dict)
171+
processed_stack_frames: list[str] = field(default_factory=list)
172+
173+
162174
class HttpResolverLocal(ApiGatewayResolver):
163175
"""
164176
ASGI-compatible HTTP resolver.
@@ -204,6 +216,8 @@ def __init__(
204216
strip_prefixes: list[str | Any] | None = None,
205217
enable_validation: bool = False,
206218
):
219+
self._startup_state = _RequestState()
220+
self._request_state: ContextVar[_RequestState | None] = ContextVar("local_http_request", default=None)
207221
super().__init__(
208222
proxy_type=ProxyEventType.APIGatewayProxyEvent, # Use REST API format internally
209223
cors=cors,
@@ -212,7 +226,46 @@ def __init__(
212226
strip_prefixes=strip_prefixes,
213227
enable_validation=enable_validation,
214228
)
215-
self._is_async_mode = False
229+
230+
@property
231+
def _state(self) -> _RequestState:
232+
return self._request_state.get() or self._startup_state
233+
234+
# Powertools declares these as mutable attributes. Properties preserve that
235+
# interface while directing each task to its own state. asyncio.to_thread
236+
# propagates the ContextVar, so middleware sees the same request dictionary.
237+
@property
238+
def current_event(self) -> BaseProxyEvent:
239+
# Preserve the inherited synchronous resolve() path outside ASGI calls.
240+
return self._state.event or BaseRouter.current_event
241+
242+
@current_event.setter
243+
def current_event(self, value: BaseProxyEvent) -> None:
244+
self._state.event = value
245+
246+
@property
247+
def lambda_context(self) -> Any:
248+
return self._state.lambda_context or BaseRouter.lambda_context
249+
250+
@lambda_context.setter
251+
def lambda_context(self, value: Any) -> None:
252+
self._state.lambda_context = value
253+
254+
@property
255+
def context(self) -> dict:
256+
return self._state.context
257+
258+
@context.setter
259+
def context(self, value: dict) -> None:
260+
self._state.context = value
261+
262+
@property
263+
def processed_stack_frames(self) -> list[str]:
264+
return self._state.processed_stack_frames
265+
266+
@processed_stack_frames.setter
267+
def processed_stack_frames(self, value: list[str]) -> None:
268+
self._state.processed_stack_frames = value
216269

217270
def _to_proxy_event(self, event: dict) -> BaseProxyEvent:
218271
"""Convert event dict to HttpProxyEvent."""
@@ -234,7 +287,7 @@ async def _resolve_async(self) -> dict: # type: ignore[override]
234287
response_builder = await super()._resolve_async()
235288
return response_builder.build(self.current_event, self._cors)
236289

237-
async def asgi_handler(self, scope: dict, receive: Callable, send: Callable) -> None:
290+
async def asgi_handler(self, scope: MutableMapping[str, Any], receive: Callable, send: Callable) -> None:
238291
"""
239292
ASGI interface - allows running with uvicorn/hypercorn/etc.
240293
@@ -274,25 +327,27 @@ async def asgi_handler(self, scope: dict, receive: Callable, send: Callable) ->
274327
# Create mock Lambda context
275328
context: Any = MockLambdaContext()
276329

277-
# Set up resolver state (similar to resolve())
278-
BaseRouter.current_event = self._to_proxy_event(event._data)
279-
BaseRouter.lambda_context = context
280-
281-
self._is_async_mode = True
282-
330+
# Never write BaseRouter's class attributes: another ASGI request may
331+
# enter while validation or the handler is awaiting I/O.
332+
state = _RequestState(
333+
event=self._to_proxy_event(event._data),
334+
lambda_context=context,
335+
context=self._startup_state.context.copy(),
336+
)
337+
token = self._request_state.set(state)
283338
try:
284-
# Use async resolve
285339
response = await self._resolve_async()
286340
finally:
287-
self._is_async_mode = False
288-
self.clear_context()
341+
# Reset only this task's binding. Middleware threads may still be
342+
# unwinding after cancellation and retain their request's state.
343+
self._request_state.reset(token)
289344

290345
# Send HTTP response
291346
await self._send_response(send, response)
292347

293348
async def __call__( # type: ignore[override]
294349
self,
295-
scope: dict,
350+
scope: MutableMapping[str, Any],
296351
receive: Callable,
297352
send: Callable,
298353
) -> None:

‎tests/functional/event_handler/_pydantic/test_http_resolver_pydantic.py‎

Lines changed: 82 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -365,3 +365,85 @@ def create_user(user: UserModel) -> UserResponse:
365365
# THEN schema includes 422 response
366366
post_operation = schema.paths["/users"].post
367367
assert 422 in post_operation.responses
368+
369+
370+
async def _post_concurrent_request(app, name: str | None):
371+
payload = {} if name is None else {"name": name, "age": 30}
372+
scope = {
373+
"type": "http",
374+
"method": "POST",
375+
"path": "/concurrent",
376+
"headers": [(b"content-type", b"application/json"), (b"x-request-name", (name or "invalid").encode())],
377+
"query_string": b"",
378+
}
379+
send, captured = make_asgi_send()
380+
await asyncio.wait_for(app(scope, make_asgi_receive(json.dumps(payload).encode()), send), timeout=5)
381+
return captured["status_code"], json.loads(captured["body"])
382+
383+
384+
def test_concurrent_asgi_validation_preserves_each_request_body():
385+
# GIVEN one local application using native request validation
386+
app = HttpResolverLocal(enable_validation=True)
387+
388+
@app.post("/concurrent")
389+
async def echo(user: UserModel) -> dict:
390+
await asyncio.sleep(0)
391+
return {"name": user.name}
392+
393+
async def scenario():
394+
# WHEN distinct bodies are submitted concurrently
395+
return await asyncio.gather(*(_post_concurrent_request(app, str(i)) for i in range(6)))
396+
397+
# THEN each caller receives its own input
398+
assert asyncio.run(scenario()) == [(200, {"name": str(i)}) for i in range(6)]
399+
400+
401+
@pytest.mark.parametrize("interruption", ["invalid", "cancelled"])
402+
def test_interrupted_asgi_request_does_not_clear_another_requests_state(interruption):
403+
async def scenario():
404+
# GIVEN an active request that reads its context after awaiting I/O
405+
app = HttpResolverLocal(enable_validation=True)
406+
entered = {name: asyncio.Event() for name in ("first", "second", "later")}
407+
release = {name: asyncio.Event() for name in entered}
408+
pending = []
409+
410+
@app.post("/concurrent")
411+
async def echo(user: UserModel) -> dict:
412+
app.append_context(name=user.name)
413+
entered[user.name].set()
414+
await release[user.name].wait()
415+
return {"name": app.context["name"], "header": app.current_event.headers["x-request-name"]}
416+
417+
async def start(name):
418+
task = asyncio.create_task(_post_concurrent_request(app, name))
419+
pending.append(task)
420+
await asyncio.wait_for(entered[name].wait(), timeout=5)
421+
return task
422+
423+
try:
424+
first = await start("first")
425+
# WHEN another request fails validation or an overlapping request is cancelled
426+
if interruption == "invalid":
427+
status, _ = await _post_concurrent_request(app, None)
428+
assert status == 422
429+
survivor, name = first, "first"
430+
else:
431+
survivor, name = await start("second"), "second"
432+
first.cancel()
433+
with pytest.raises(asyncio.CancelledError):
434+
await first
435+
436+
# THEN the surviving request retains both context and headers
437+
release[name].set()
438+
assert await survivor == (200, {"name": name, "header": name})
439+
release["later"].set()
440+
assert await _post_concurrent_request(app, "later") == (200, {"name": "later", "header": "later"})
441+
finally:
442+
for event in release.values():
443+
event.set()
444+
for task in pending:
445+
if not task.done():
446+
task.cancel()
447+
await asyncio.gather(*pending, return_exceptions=True)
448+
449+
asyncio.run(scenario())

0 commit comments

Comments
 (0)