From 5163fe37c545ae27bad6fe603f563fbe80893991 Mon Sep 17 00:00:00 2001 From: Sasha Mitchell Date: Fri, 9 Oct 2026 16:24:41 +0700 Subject: [PATCH] fix: signal normal completion on REST SSE streams REST sendMessageStreaming and subscribeToTask dropped the HTTP client's normal completion callback. JSON-RPC already turns that callback into a null terminal signal. Parse errors now use the same one-shot callback so a later completion cannot replace them with success. Fixes #1194 --- .../client/transport/rest/RestTransport.java | 6 ++- .../transport/rest/sse/SSEEventListener.java | 14 ++++--- .../transport/rest/RestTransportTest.java | 42 +++++++++++++++++++ .../rest/sse/SSEEventListenerTest.java | 40 ++++++++++++++++++ 4 files changed, 94 insertions(+), 8 deletions(-) diff --git a/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/RestTransport.java b/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/RestTransport.java index 1f6535250..68449dcbb 100644 --- a/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/RestTransport.java +++ b/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/RestTransport.java @@ -121,7 +121,8 @@ public void sendMessageStreaming(MessageSendParams messageSendParams, Consumer sseEventListener.onMessage(event, ref.get()), throwable -> sseEventListener.onError(throwable, ref.get()), () -> { - // We don't need to do anything special on completion + // Signal normal stream completion to error handler (null error means success) + sseEventListener.onComplete(); })); } catch (IOException e) { throw new A2AClientException("Failed to send streaming message request: " + e, e); @@ -381,7 +382,8 @@ public void subscribeToTask(TaskIdParams request, Consumer e event -> sseEventListener.onMessage(event, ref.get()), throwable -> sseEventListener.onError(throwable, ref.get()), () -> { - // We don't need to do anything special on completion + // Signal normal stream completion to error handler (null error means success) + sseEventListener.onComplete(); })); } catch (IOException e) { throw new A2AClientException("Failed to send streaming message request: " + e, e); diff --git a/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/sse/SSEEventListener.java b/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/sse/SSEEventListener.java index d2507c64e..783a50e8e 100644 --- a/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/sse/SSEEventListener.java +++ b/client/transport/rest/src/main/java/org/a2aproject/sdk/client/transport/rest/sse/SSEEventListener.java @@ -35,12 +35,15 @@ public void onMessage(ServerSentEvent event, @Nullable Future completableF JsonFormat.parser().merge(event.data(), builder); parseAndHandleMessage(builder.build(), completableFuture); } catch (InvalidProtocolBufferException e) { - if (getErrorHandler() != null) { - getErrorHandler().accept(RestErrorMapper.mapRestError(event.data(), 500)); - } + signalTerminal(RestErrorMapper.mapRestError(event.data(), 500)); } } + public void onComplete() { + LOGGER.fine("SSEEventListener.onComplete() called - signaling successful stream completion"); + signalTerminal(null); + } + /** * Parses a StreamResponse protobuf message and delegates to the base class for event handling. * @@ -60,9 +63,8 @@ private void parseAndHandleMessage(StreamResponse response, @Nullable Future { LOGGER.warning("Invalid stream response " + response.getPayloadCase()); - if (getErrorHandler() != null) { - getErrorHandler().accept(new IllegalStateException("Invalid stream response from server: " + response.getPayloadCase())); - } + signalTerminal(new IllegalStateException( + "Invalid stream response from server: " + response.getPayloadCase())); return; } } diff --git a/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/RestTransportTest.java b/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/RestTransportTest.java index 05d1b7b12..884ccad00 100644 --- a/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/RestTransportTest.java +++ b/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/RestTransportTest.java @@ -284,6 +284,48 @@ public void testSendMessageStreaming() throws Exception { assertEquals("2", task.id()); } + /** + * A non-final SSE body still has to report normal completion when the connection ends. + * JSON-RPC already does this. REST used to drop the callback. + */ + @Test + public void testSendMessageStreamingSignalsNormalCompletion() throws Exception { + String streamResponseBody = "event: message\n" + + "data: {\"task\":{\"id\":\"2\",\"contextId\":\"context-open\",\"status\":{\"state\":\"TASK_STATE_SUBMITTED\"}}}\n\n"; + this.server.when( + request() + .withMethod("POST") + .withPath("/message:stream") + ) + .respond( + response() + .withStatusCode(200) + .withHeader("Content-Type", "text/event-stream") + .withBody(streamResponseBody) + ); + + RestTransport client = new RestTransport(CARD); + Message message = Message.builder() + .role(Message.Role.ROLE_USER) + .parts(Collections.singletonList(new TextPart("still working"))) + .contextId("context-open") + .messageId("message-open") + .build(); + MessageSendParams params = MessageSendParams.builder() + .message(message) + .build(); + + AtomicReference terminal = new AtomicReference<>(new RuntimeException("unset")); + CountDownLatch latch = new CountDownLatch(1); + client.sendMessageStreaming(params, event -> { }, error -> { + terminal.set(error); + latch.countDown(); + }, null); + + assertTrue(latch.await(10, TimeUnit.SECONDS)); + assertNull(terminal.get()); + } + /** * Test of CreateTaskPushNotificationConfiguration method, of class JSONRestTransport. */ diff --git a/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/sse/SSEEventListenerTest.java b/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/sse/SSEEventListenerTest.java index 65ada80b1..48def5fc1 100644 --- a/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/sse/SSEEventListenerTest.java +++ b/client/transport/rest/src/test/java/org/a2aproject/sdk/client/transport/rest/sse/SSEEventListenerTest.java @@ -2,11 +2,13 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertTrue; import java.util.concurrent.Future; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicReference; import org.a2aproject.sdk.client.http.ServerSentEvent; @@ -283,6 +285,44 @@ public void testOnMessageWithInvalidPayloadCaseCallsErrorHandler() { assertTrue(receivedError.get().getMessage().contains("Invalid stream response")); } + @Test + public void testOnCompleteSignalsNormalCompletion() { + AtomicReference received = new AtomicReference<>(new RuntimeException("unset")); + AtomicBoolean called = new AtomicBoolean(false); + SSEEventListener listener = new SSEEventListener( + event -> {}, + error -> { + called.set(true); + received.set(error); + } + ); + + listener.onComplete(); + + assertTrue(called.get()); + assertNull(received.get()); + } + + @Test + public void testOnErrorThenOnCompleteKeepsTheError() { + AtomicInteger calls = new AtomicInteger(); + AtomicReference received = new AtomicReference<>(); + IllegalStateException boom = new IllegalStateException("first"); + SSEEventListener listener = new SSEEventListener( + event -> {}, + error -> { + calls.incrementAndGet(); + received.set(error); + } + ); + + listener.onError(boom, null); + listener.onComplete(); + + assertEquals(1, calls.get()); + assertEquals(boom, received.get()); + } + @Test public void testOnErrorCallsErrorHandler() { // Set up error handler