From b239dec5736618fbef45499850bf99de93c08256 Mon Sep 17 00:00:00 2001 From: Anas Khan <83116240+anxkhn@users.noreply.github.com> Date: Fri, 28 Aug 2026 02:26:13 +0530 Subject: [PATCH] fix: do not treat TASK_STATE_UNSPECIFIED as a terminal event Signed-off-by: Anas Khan <83116240+anxkhn@users.noreply.github.com> --- src/a2a/server/events/event_consumer.py | 1 - tests/server/events/test_event_consumer.py | 37 ++++++++++++++++++++++ 2 files changed, 37 insertions(+), 1 deletion(-) diff --git a/src/a2a/server/events/event_consumer.py b/src/a2a/server/events/event_consumer.py index 8414e2d17..075439205 100644 --- a/src/a2a/server/events/event_consumer.py +++ b/src/a2a/server/events/event_consumer.py @@ -75,7 +75,6 @@ async def consume_all(self) -> AsyncGenerator[Event]: TaskState.TASK_STATE_CANCELED, TaskState.TASK_STATE_FAILED, TaskState.TASK_STATE_REJECTED, - TaskState.TASK_STATE_UNSPECIFIED, TaskState.TASK_STATE_INPUT_REQUIRED, ) ) diff --git a/tests/server/events/test_event_consumer.py b/tests/server/events/test_event_consumer.py index fb0f878a1..81ae50c69 100644 --- a/tests/server/events/test_event_consumer.py +++ b/tests/server/events/test_event_consumer.py @@ -103,6 +103,43 @@ async def mock_dequeue() -> Any: assert mock_event_queue.task_done.call_count == 3 +@pytest.mark.asyncio +async def test_consume_all_does_not_stop_on_unspecified_state( + event_consumer: MagicMock, + mock_event_queue: MagicMock, +): + events: list[Any] = [ + Task(id='task_123', context_id='session-xyz'), + TaskStatusUpdateEvent( + task_id='task_123', + context_id='session-xyz', + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ), + TaskStatusUpdateEvent( + task_id='task_123', + context_id='session-xyz', + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ), + ] + cursor = 0 + + async def mock_dequeue() -> Any: + nonlocal cursor + if cursor < len(events): + event = events[cursor] + cursor += 1 + return event + mock_event_queue.is_closed.return_value = True + raise asyncio.QueueEmpty() + + mock_event_queue.dequeue_event = mock_dequeue + consumed_events: list[Any] = [] + async for event in event_consumer.consume_all(): + consumed_events.append(event) + assert consumed_events == events + assert mock_event_queue.task_done.call_count == 3 + + @pytest.mark.asyncio async def test_consume_until_message( event_consumer: MagicMock,