Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions src/a2a/server/agent_execution/active_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
*,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
141 changes: 141 additions & 0 deletions tests/server/request_handlers/test_default_request_handler_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading