From 85283a9d73418aa6fced1fa0f7b769e2edd2c092 Mon Sep 17 00:00:00 2001 From: Dang Zitou Date: Thu, 13 Aug 2026 02:48:08 +0800 Subject: [PATCH] fix(server-common): keep task streams open after messages This fixes #1037 --- .../sdk/server/events/EventConsumer.java | 6 ++-- .../sdk/server/events/EventConsumerTest.java | 35 +++++++++++++++++++ 2 files changed, 38 insertions(+), 3 deletions(-) diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/events/EventConsumer.java b/server-common/src/main/java/org/a2aproject/sdk/server/events/EventConsumer.java index 531421aca..3d8418c27 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/events/EventConsumer.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/events/EventConsumer.java @@ -228,9 +228,9 @@ public Flow.Publisher consumeAll() { boolean isFinalEvent = false; if (event instanceof TaskStatusUpdateEvent tue && tue.isFinal()) { isFinalEvent = true; - } else if (event instanceof Message) { - // Per A2A spec §3.1.2 (Send Streaming Message): a Message is the - // complete response — the stream must close after delivering it. + } else if (event instanceof Message && lastSeenTaskState == null) { + // A stateless Message is the complete response. Messages emitted + // after a task has started are intermediate task stream events. isFinalEvent = true; } else if (event instanceof Task task) { isFinalEvent = isStreamTerminatingTask(task); diff --git a/server-common/src/test/java/org/a2aproject/sdk/server/events/EventConsumerTest.java b/server-common/src/test/java/org/a2aproject/sdk/server/events/EventConsumerTest.java index 54882a078..396e66a87 100644 --- a/server-common/src/test/java/org/a2aproject/sdk/server/events/EventConsumerTest.java +++ b/server-common/src/test/java/org/a2aproject/sdk/server/events/EventConsumerTest.java @@ -236,6 +236,41 @@ public void testConsumeMessageEvents() throws Exception { assertSame(message, receivedEvents.get(0)); } + @Test + public void testConsumeTaskMessageBeforeInputRequired() throws Exception { + Task workingTask = Task.builder() + .id(TASK_ID) + .contextId("session-xyz") + .status(new TaskStatus(TaskState.TASK_STATE_WORKING)) + .build(); + Message message = fromJson(MESSAGE_PAYLOAD, Message.class); + TaskStatusUpdateEvent inputRequiredEvent = TaskStatusUpdateEvent.builder() + .taskId(TASK_ID) + .contextId("session-xyz") + .status(new TaskStatus(TaskState.TASK_STATE_INPUT_REQUIRED)) + .build(); + TaskStatusUpdateEvent completedEvent = TaskStatusUpdateEvent.builder() + .taskId(TASK_ID) + .contextId("session-xyz") + .status(new TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .build(); + List events = List.of(workingTask, message, inputRequiredEvent, completedEvent); + + for (Event event : events) { + eventQueue.enqueueEvent(event); + } + + List receivedEvents = new ArrayList<>(); + AtomicReference error = new AtomicReference<>(); + eventConsumer.consumeAll().subscribe(getSubscriber(receivedEvents, error)); + + assertNull(error.get()); + assertEquals(events.size(), receivedEvents.size()); + for (int i = 0; i < events.size(); i++) { + assertSame(events.get(i), receivedEvents.get(i)); + } + } + @Test public void testBufferFlushDelayMsDefaultsTo150() { String original = System.getProperty("a2a.eventconsumer.bufferFlushDelayMs");