From d9c5fcf2c4bc6f994dd2be7059467e85dd67fa19 Mon Sep 17 00:00:00 2001 From: Tejas Kashinath Date: Mon, 5 Oct 2026 15:27:22 -0400 Subject: [PATCH] fix(server): run follow-up messages in the sender's contextvars ActiveTask creates its producer once, during the first request for a task, so every later message on the task ran AgentExecutor.execute() with the first request's contextvars. Capture the caller's context in enqueue_request() and run execute() in it. Fixes #1316 --- src/a2a/server/agent_execution/active_task.py | 27 ++++-- .../test_default_request_handler_v2.py | 85 +++++++++++++++++++ 2 files changed, 105 insertions(+), 7 deletions(-) diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index 8cb648bae..a2f174367 100644 --- a/src/a2a/server/agent_execution/active_task.py +++ b/src/a2a/server/agent_execution/active_task.py @@ -36,6 +36,7 @@ from __future__ import annotations import asyncio +import contextvars import logging import uuid @@ -468,9 +469,9 @@ def __init__( # noqa: PLR0913 self._reference_count = 0 # Queue for incoming requests - self._request_queue: AsyncQueue[tuple[RequestContext, uuid.UUID]] = ( - create_async_queue() - ) + self._request_queue: AsyncQueue[ + tuple[RequestContext, uuid.UUID, contextvars.Context] + ] = create_async_queue() @property def task_id(self) -> str: @@ -480,9 +481,16 @@ def task_id(self) -> str: async def enqueue_request( self, request_context: RequestContext ) -> uuid.UUID: - """Enqueues a request for the active task to process.""" + """Enqueues a request for the active task to process. + + The caller's contextvars are captured so the producer runs + `AgentExecutor.execute` in the context of the request that sent the + message, not the request that started the producer. + """ request_id = uuid.uuid4() - await self._request_queue.put((request_context, request_id)) + await self._request_queue.put( + (request_context, request_id, contextvars.copy_context()) + ) return request_id async def start( @@ -573,6 +581,7 @@ async def _run_producer(self) -> None: ( request_context, request_id, + sender_context, ) = await self._request_queue.get() await self._request_lock.acquire() # TODO: Should we create task manager every time? @@ -607,8 +616,12 @@ async def _run_producer(self) -> None: _RequestStarted(request_id, request_context), ) ) - await self._agent_executor.execute( - request_context, self._event_queue_agent + # Awaiting the child task propagates producer cancellation. + await sender_context.run( + asyncio.create_task, + self._agent_executor.execute( + request_context, self._event_queue_agent + ), ) logger.debug( 'Producer[%s]: Execution finished successfully', diff --git a/tests/server/request_handlers/test_default_request_handler_v2.py b/tests/server/request_handlers/test_default_request_handler_v2.py index 1472a0bc6..52c5e9c4b 100644 --- a/tests/server/request_handlers/test_default_request_handler_v2.py +++ b/tests/server/request_handlers/test_default_request_handler_v2.py @@ -1,4 +1,5 @@ import asyncio +import contextvars import logging import time import uuid @@ -2443,3 +2444,87 @@ async def test_cancel_of_input_required_task_cannot_be_undone(): assert stored.status.state == TaskState.TASK_STATE_CANCELED assert agent.execute_calls == 1 await handler.aclose() + + +_REQUEST_TAG: contextvars.ContextVar[str] = contextvars.ContextVar( + 'request_tag', default='unset' +) + + +class _InputRequiredThenCompleteAgent(AgentExecutor): + """Asks for input on the first message of a task and completes on the next.""" + + def __init__(self) -> None: + self.seen_tags: list[str] = [] + self.task: Task | None = None + + async def execute( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + self.seen_tags.append(_REQUEST_TAG.get()) + if context.current_task: + updater = TaskUpdater( + event_queue, context.task_id, context.context_id + ) + await updater.complete() + return + self.task = new_task_from_user_message(context.message) + await event_queue.enqueue_event(self.task) + updater = TaskUpdater(event_queue, self.task.id, self.task.context_id) + await updater.update_status(TaskState.TASK_STATE_INPUT_REQUIRED) + + async def cancel( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + pass + + +@pytest.mark.timeout(10) +@pytest.mark.asyncio +@pytest.mark.parametrize('streaming', [False, True]) +async def test_follow_up_message_runs_in_sender_contextvars(streaming): + """A follow-up message on a live task must see its own request's contextvars.""" + agent = _InputRequiredThenCompleteAgent() + handler = DefaultRequestHandlerV2( + agent_executor=agent, + task_store=InMemoryTaskStore(), + agent_card=create_default_agent_card(), + ) + + async def send(tag: str, message: Message) -> None: + _REQUEST_TAG.set(tag) + params = SendMessageRequest(message=message) + if streaming: + async for _ in handler.on_message_send_stream( + params, create_server_call_context() + ): + pass + else: + await handler.on_message_send(params, create_server_call_context()) + + await asyncio.create_task( + send( + 'request-1', + Message( + message_id='msg-1', + role=Role.ROLE_USER, + parts=[Part(text='book a flight')], + ), + ) + ) + assert agent.task is not None + await asyncio.create_task( + send( + 'request-2', + Message( + message_id='msg-2', + role=Role.ROLE_USER, + parts=[Part(text='Friday')], + task_id=agent.task.id, + context_id=agent.task.context_id, + ), + ) + ) + + assert agent.seen_tags == ['request-1', 'request-2'] + await handler.aclose()