diff --git a/src/bedrock_agentcore/runtime/a2a.py b/src/bedrock_agentcore/runtime/a2a.py index bd0d6261..1c082e2d 100644 --- a/src/bedrock_agentcore/runtime/a2a.py +++ b/src/bedrock_agentcore/runtime/a2a.py @@ -4,6 +4,8 @@ extraction, health checks, and Docker host detection. """ +import asyncio +import contextvars import logging import uuid from importlib import import_module @@ -34,6 +36,8 @@ A2A_CONTRACT_PORT = 9000 A2A_PORT_ENV = "A2A_PORT" +_CONTEXTVARS_STATE_KEY = "_bedrock_agentcore_contextvars" + def _check_a2a_sdk() -> None: """Raise ImportError with install instructions if a2a-sdk is missing.""" @@ -247,6 +251,37 @@ def build(self, request: Any) -> Any: pass +class _ContextvarsSnapshotRequestContextBuilder: + """Stores a copy of the sending request's contextvars on each RequestContext. + + a2a-sdk 1.x runs every message of a task in one producer task created by the + first request, so the executor would otherwise see that request's contextvars. + Remove with ``_PerMessageContextvarsExecutor`` once a2a-python#1316 is fixed. + """ + + def __init__(self, inner: Any) -> None: + self._inner = inner + + async def build(self, *args: Any, **kwargs: Any) -> Any: + request_context = await self._inner.build(*args, **kwargs) + request_context.call_context.state[_CONTEXTVARS_STATE_KEY] = contextvars.copy_context() + return request_context + + +class _PerMessageContextvarsExecutor: + """Runs the wrapped executor in the contextvars of the request that sent the message.""" + + def __init__(self, inner: Any) -> None: + self._inner = inner + + async def execute(self, context: Any, event_queue: Any) -> None: + snapshot = context.call_context.state[_CONTEXTVARS_STATE_KEY] + await snapshot.run(asyncio.create_task, self._inner.execute(context, event_queue)) + + async def cancel(self, context: Any, event_queue: Any) -> None: + await self._inner.cancel(context, event_queue) + + def build_a2a_app( executor: Any, agent_card: Any = None, @@ -300,12 +335,16 @@ def build_a2a_app( request_handler_type: Any = DefaultRequestHandler if is_a2a_v1: + from a2a.server.agent_execution import SimpleRequestContextBuilder from a2a.server.routes import create_agent_card_routes, create_jsonrpc_routes http_handler = request_handler_type( - agent_executor=executor, + agent_executor=_PerMessageContextvarsExecutor(executor), task_store=task_store, agent_card=agent_card, + request_context_builder=_ContextvarsSnapshotRequestContextBuilder( + SimpleRequestContextBuilder(should_populate_referred_tasks=False, task_store=task_store) + ), ) routes = create_agent_card_routes(agent_card) routes.extend( diff --git a/tests/bedrock_agentcore/runtime/test_a2a.py b/tests/bedrock_agentcore/runtime/test_a2a.py index 3bdc04a2..7c5afea7 100644 --- a/tests/bedrock_agentcore/runtime/test_a2a.py +++ b/tests/bedrock_agentcore/runtime/test_a2a.py @@ -53,6 +53,42 @@ async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None pass +_REQUEST_TAG: contextvars.ContextVar[str] = contextvars.ContextVar("request_tag", default="unset") + + +class _RequestTagMiddleware: + def __init__(self, app): + self.app = app + + async def __call__(self, scope, receive, send): + token = _REQUEST_TAG.set(dict(scope.get("headers", [])).get(b"x-request-tag", b"unset").decode()) + try: + await self.app(scope, receive, send) + finally: + _REQUEST_TAG.reset(token) + + +class _InputRequiredExecutor(AgentExecutor): + """Asks for input on the first message of a task and completes on the follow-up.""" + + def __init__(self): + self.seen = [] + + async def execute(self, context: RequestContext, event_queue: EventQueue) -> None: + token = BedrockAgentCoreContext.get_workload_access_token() + self.seen.append((context.task_id, context.context_id, _REQUEST_TAG.get(), token)) + task = context.current_task or _new_task(context.message) + updater = TaskUpdater(event_queue, task.id, task.context_id) + if context.current_task: + await updater.complete() + else: + await event_queue.enqueue_event(task) + await updater.requires_input() + + async def cancel(self, context: RequestContext, event_queue: EventQueue) -> None: + pass + + def _make_agent_card() -> AgentCard: card_kwargs = { "name": "test-agent", @@ -116,22 +152,23 @@ def _jsonrpc_request(method: str, params: dict | None = None) -> dict: return body -def _send_message_params(text: str = "hello") -> dict: +def _send_message_params(text: str = "hello", task_id: str | None = None, context_id: str | None = None) -> dict: if not IS_A2A_V1: - return { - "message": { - "message_id": str(uuid.uuid4()), - "role": "user", - "parts": [{"kind": "text", "text": text}], - } + message = { + "message_id": str(uuid.uuid4()), + "role": "user", + "parts": [{"kind": "text", "text": text}], } - return { - "message": { + else: + message = { "messageId": str(uuid.uuid4()), "role": "ROLE_USER", "parts": [{"text": text}], } - } + if task_id: + message["taskId"] = task_id + message["contextId"] = context_id + return {"message": message} class TestBuildA2AApp: @@ -194,6 +231,31 @@ def test_message_send_executes_and_returns_completed_task(self): assert task["status"]["state"] == _completed_state() assert task["artifacts"][0]["parts"][0]["text"] == "echo: unit-test" + @pytest.mark.parametrize( + "method", + [_method("SendMessage", "message/send"), _method("SendStreamingMessage", "message/stream")], + ) + def test_follow_up_message_runs_in_its_own_request_contextvars(self, method): + """Regression for #690: a2a-sdk 1.x reuses the first request's producer task for every message of a task.""" + executor = _InputRequiredExecutor() + app = build_a2a_app(executor, _make_agent_card()) + app.add_middleware(_RequestTagMiddleware) + + def send(tag, **ids): + resp = client.post( + "/", + json=_jsonrpc_request(method, _send_message_params(tag, **ids)), + headers={"x-request-tag": tag, "WorkloadAccessToken": f"wat-{tag}"}, + ) + assert resp.status_code == 200 + + with _test_client(app) as client: + send("turn-1") + task_id, context_id = executor.seen[0][:2] + send("turn-2", task_id=task_id, context_id=context_id) + + assert [row[2:] for row in executor.seen] == [("turn-1", "wat-turn-1"), ("turn-2", "wat-turn-2")] + @pytest.mark.skipif(not IS_A2A_V1, reason="v0.3 compatibility routes are provided by a2a-sdk v1") def test_v03_message_send_remains_compatible(self): app = build_a2a_app(_EchoExecutor(), _make_agent_card())