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
Original file line number Diff line number Diff line change
Expand Up @@ -121,7 +121,8 @@ public void sendMessageStreaming(MessageSendParams messageSendParams, Consumer<S
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);
Expand Down Expand Up @@ -381,7 +382,8 @@ public void subscribeToTask(TaskIdParams request, Consumer<StreamingEventKind> 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);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,12 +35,15 @@ public void onMessage(ServerSentEvent event, @Nullable Future<Void> 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.
*
Expand All @@ -60,9 +63,8 @@ private void parseAndHandleMessage(StreamResponse response, @Nullable Future<Voi
event = ProtoUtils.FromProto.taskArtifactUpdateEvent(response.getArtifactUpdate());
default -> {
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;
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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<Throwable> 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.
*/
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -283,6 +285,44 @@ public void testOnMessageWithInvalidPayloadCaseCallsErrorHandler() {
assertTrue(receivedError.get().getMessage().contains("Invalid stream response"));
}

@Test
public void testOnCompleteSignalsNormalCompletion() {
AtomicReference<Throwable> 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<Throwable> 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
Expand Down
Loading