From 9a3e78a4be4b8800976b79ed35f82934fa50f34f Mon Sep 17 00:00:00 2001 From: Sasha Mitchell Date: Fri, 9 Oct 2026 13:42:26 +0700 Subject: [PATCH] fix(server): resubscribe a paused task from the shared stream A replica that finished an interrupted turn keeps the task registered. Tapping that entry returned the old snapshot and never saw the turn running elsewhere. The local tap now stays for a turn this replica is executing. --- src/a2a/server/agent_execution/active_task.py | 8 + .../default_request_handler_v2.py | 6 +- .../test_default_request_handler_v2.py | 141 ++++++++++++++++++ 3 files changed, 153 insertions(+), 2 deletions(-) diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index a2f174367..c51dc2e47 100644 --- a/src/a2a/server/agent_execution/active_task.py +++ b/src/a2a/server/agent_execution/active_task.py @@ -688,6 +688,14 @@ async def _run_consumer(self) -> None: logger.debug('Consumer[%s]: Finishing', self._task_id) await self._maybe_cleanup() + def request_in_flight(self) -> bool: + """True while this replica is executing a turn for the task. + + The lock is held from the start of a request until its completion + event. An interrupted turn releases it and leaves the task registered. + """ + return self._request_lock.locked() + async def subscribe( self, *, diff --git a/src/a2a/server/request_handlers/default_request_handler_v2.py b/src/a2a/server/request_handlers/default_request_handler_v2.py index b135f529f..5ea40c433 100644 --- a/src/a2a/server/request_handlers/default_request_handler_v2.py +++ b/src/a2a/server/request_handlers/default_request_handler_v2.py @@ -581,10 +581,12 @@ async def on_subscribe_to_task( # noqa: D102 yield event return - # Shared-stream mode. Fast path: this replica runs the agent -> tap it. + # Shared-stream mode. Tap this replica only while it is executing. + # A registry entry also outlives an interrupted turn, and that cached + # snapshot is not the live task. stream = self._event_stream local = await self._active_task_registry.get(task_id) - if local is not None: + if local is not None and local.request_in_flight(): async for event in local.subscribe(include_initial_task=True): yield event return 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 b2d25f95a..274612203 100644 --- a/tests/server/request_handlers/test_default_request_handler_v2.py +++ b/tests/server/request_handlers/test_default_request_handler_v2.py @@ -2591,3 +2591,144 @@ def test_init_does_not_warn_for_complete_agent_cards( extended_agent_card=_complete_agent_card(), ) assert caplog.records == [] + + +class _ScriptedLocal: + """Stand-in for an ActiveTask that is either executing or parked.""" + + def __init__(self, in_flight: bool, events: list[Task]) -> None: + self._in_flight = in_flight + self._events = events + self.subscribed = False + + def request_in_flight(self) -> bool: + return self._in_flight + + async def subscribe(self, **kwargs: object): + del kwargs + self.subscribed = True + for event in self._events: + yield event + if not self._in_flight: + await asyncio.Event().wait() + + +class _ScriptedStream: + """Yields one completed update and counts how often it is tailed.""" + + def __init__(self, event: TaskStatusUpdateEvent) -> None: + self._event = event + self.subscribed = 0 + + async def publish(self, task_id: str, event: object) -> None: + del task_id, event + + def subscribe(self, task_id: str, *, after: object): + del task_id, after + self.subscribed += 1 + return self._events() + + async def _events(self): + from a2a.server.cluster.event_stream import VersionedEvent + from a2a.server.cluster.version import TaskVersion + + yield VersionedEvent(event=self._event, version=TaskVersion(1)) + + async def destroy(self, task_id: str) -> None: + del task_id + + +async def _collect_subscribe(handler, task_id: str, context): + events = [] + async for event in handler.on_subscribe_to_task( + SubscribeToTaskRequest(id=task_id), context + ): + events.append(event) + return events + + +@pytest.mark.timeout(10) +@pytest.mark.asyncio +async def test_paused_replica_resubscribe_follows_the_shared_stream(): + """A parked local task must not hide a turn running on another replica.""" + alice = _ctx('alice') + store = InMemoryTaskStore() + current = create_sample_task( + 'task-1', TaskState.TASK_STATE_WORKING, context_id='ctx-1' + ) + await store.save(current, alice) + completed = TaskStatusUpdateEvent( + task_id='task-1', + context_id='ctx-1', + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + stream = _ScriptedStream(completed) + stale = create_sample_task( + 'task-1', TaskState.TASK_STATE_INPUT_REQUIRED, context_id='ctx-1' + ) + local = _ScriptedLocal(in_flight=False, events=[stale]) + handler = DefaultRequestHandlerV2( + agent_executor=MockAgentExecutor(), + task_store=store, + agent_card=create_default_agent_card(), + event_stream=stream, + ) + + async def get_local(task_id: str): + del task_id + return local + + handler._active_task_registry.get = get_local # type: ignore[method-assign] + + events = await asyncio.wait_for( + _collect_subscribe(handler, 'task-1', alice), timeout=2 + ) + + assert [event.status.state for event in events] == [ + TaskState.TASK_STATE_WORKING, + TaskState.TASK_STATE_COMPLETED, + ] + assert local.subscribed is False + assert stream.subscribed == 1 + + +@pytest.mark.timeout(10) +@pytest.mark.asyncio +async def test_in_flight_replica_resubscribe_stays_on_the_local_task(): + """A turn running on this replica is still tapped locally.""" + alice = _ctx('alice') + store = InMemoryTaskStore() + stored = create_sample_task( + 'task-1', TaskState.TASK_STATE_WORKING, context_id='ctx-1' + ) + await store.save(stored, alice) + completed = TaskStatusUpdateEvent( + task_id='task-1', + context_id='ctx-1', + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + stream = _ScriptedStream(completed) + live = create_sample_task( + 'task-1', TaskState.TASK_STATE_WORKING, context_id='ctx-1' + ) + local = _ScriptedLocal(in_flight=True, events=[live]) + handler = DefaultRequestHandlerV2( + agent_executor=MockAgentExecutor(), + task_store=store, + agent_card=create_default_agent_card(), + event_stream=stream, + ) + + async def get_local(task_id: str): + del task_id + return local + + handler._active_task_registry.get = get_local # type: ignore[method-assign] + + events = await asyncio.wait_for( + _collect_subscribe(handler, 'task-1', alice), timeout=2 + ) + + assert events == [live] + assert local.subscribed is True + assert stream.subscribed == 0