diff --git a/a2a/src/main/java/com/google/adk/a2a/agent/RemoteA2AAgent.java b/a2a/src/main/java/com/google/adk/a2a/agent/RemoteA2AAgent.java index d4e094710..5ea71de70 100644 --- a/a2a/src/main/java/com/google/adk/a2a/agent/RemoteA2AAgent.java +++ b/a2a/src/main/java/com/google/adk/a2a/agent/RemoteA2AAgent.java @@ -36,13 +36,13 @@ import com.google.genai.types.Part; import io.a2a.client.Client; import io.a2a.client.ClientEvent; +import io.a2a.client.MessageEvent; import io.a2a.client.TaskEvent; import io.a2a.client.TaskUpdateEvent; import io.a2a.spec.A2AClientException; import io.a2a.spec.AgentCard; import io.a2a.spec.Message; import io.a2a.spec.TaskArtifactUpdateEvent; -import io.a2a.spec.TaskState; import io.a2a.spec.TaskStatusUpdateEvent; import io.reactivex.rxjava3.core.BackpressureStrategy; import io.reactivex.rxjava3.core.Flowable; @@ -227,8 +227,7 @@ protected Flowable runAsyncImpl(InvocationContext invocationContext) { return Flowable.create( emitter -> { StreamHandler handler = - new StreamHandler( - emitter.serialize(), invocationContext, requestJson, streaming, name()); + new StreamHandler(emitter.serialize(), invocationContext, requestJson, name()); ImmutableList> consumers = ImmutableList.of(handler::handleEvent); a2aClient.sendMessage(originalMessage, consumers, handler::handleError, null); @@ -249,7 +248,6 @@ private static class StreamHandler { private final FlowableEmitter emitter; private final InvocationContext invocationContext; private final String requestJson; - private final boolean streaming; private final String agentName; private boolean done = false; private final StringBuilder textBuffer = new StringBuilder(); @@ -259,12 +257,10 @@ private static class StreamHandler { FlowableEmitter emitter, InvocationContext invocationContext, String requestJson, - boolean streaming, String agentName) { this.emitter = emitter; this.invocationContext = invocationContext; this.requestJson = requestJson; - this.streaming = streaming; this.agentName = agentName; } @@ -522,13 +518,13 @@ private Event createAggregatedEvent(Content content, @Nullable ClientEvent trigg } private static boolean isCompleted(ClientEvent event) { - TaskState executionState = TaskState.UNKNOWN; if (event instanceof TaskEvent taskEvent) { - executionState = taskEvent.getTask().getStatus().state(); - } else if (event instanceof TaskUpdateEvent updateEvent) { - executionState = updateEvent.getTask().getStatus().state(); + return taskEvent.getTask().getStatus().state().isFinal(); + } + if (event instanceof TaskUpdateEvent updateEvent) { + return updateEvent.getTask().getStatus().state().isFinal(); } - return executionState.equals(TaskState.COMPLETED); + return false; } private static ImmutableList eventParts(Event event) { diff --git a/a2a/src/test/java/com/google/adk/a2a/agent/RemoteA2AAgentTest.java b/a2a/src/test/java/com/google/adk/a2a/agent/RemoteA2AAgentTest.java index 8d6c9b062..91c104008 100644 --- a/a2a/src/test/java/com/google/adk/a2a/agent/RemoteA2AAgentTest.java +++ b/a2a/src/test/java/com/google/adk/a2a/agent/RemoteA2AAgentTest.java @@ -44,6 +44,7 @@ import com.google.genai.types.Part; import io.a2a.client.Client; import io.a2a.client.ClientEvent; +import io.a2a.client.MessageEvent; import io.a2a.client.TaskEvent; import io.a2a.client.TaskUpdateEvent; import io.a2a.spec.AgentCapabilities; @@ -761,6 +762,24 @@ private ClientEvent createFinalEvent(String text) { return createTestEvent(new TextPart(text), TaskState.COMPLETED, false, false); } + private ClientEvent createFailedEvent(String errorMessage) { + return createTestEvent(new TextPart(errorMessage), TaskState.FAILED, false, false); + } + + private ClientEvent createCanceledEvent(String message) { + return createTestEvent(new TextPart(message), TaskState.CANCELED, false, false); + } + + private ClientEvent createMessageEvent(String text) { + Message message = + new Message.Builder() + .messageId("msg-id-1") + .role(Message.Role.AGENT) + .parts(new TextPart(text)) + .build(); + return new MessageEvent(message); + } + private ClientEvent createTestEvent( io.a2a.spec.Part part, TaskState state, boolean append, boolean lastChunk) { Artifact artifact = @@ -788,6 +807,62 @@ private ClientEvent createTestEvent( return new TaskUpdateEvent(task, updateEvent); } + @Test + public void runAsync_terminatesOnFailureTaskState() { + RemoteA2AAgent agent = createAgent(); + mockStreamResponse( + consumer -> { + consumer.accept(createPartialEvent("Processing data...", true, false), agentCard); + consumer.accept(createFailedEvent("Internal Server Error"), agentCard); + }); + + List events = agent.runAsync(invocationContext).toList().blockingGet(); + + // The stream must terminate cleanly and include the final event with error info. + assertThat(events).hasSize(3); + assertText(events.get(0), "Processing data..."); + assertText(events.get(1), "Processing data..."); // aggregated + assertText(events.get(2), "Internal Server Error"); // terminal failure event + } + + @Test + public void runAsync_terminatesOnCanceledTaskState() { + RemoteA2AAgent agent = createAgent(); + mockStreamResponse( + consumer -> { + consumer.accept(createPartialEvent("Processing data...", true, false), agentCard); + consumer.accept(createCanceledEvent("Execution Canceled"), agentCard); + }); + + List events = agent.runAsync(invocationContext).toList().blockingGet(); + + // The stream must terminate cleanly and include the final event with cancellation info. + assertThat(events).hasSize(3); + assertText(events.get(0), "Processing data..."); + assertText(events.get(1), "Processing data..."); // aggregated + assertText(events.get(2), "Execution Canceled"); // terminal canceled event + } + + @Test + public void runAsync_doesNotTerminateOnMessageEvent() { + RemoteA2AAgent agent = createAgent(); + mockStreamResponse( + consumer -> { + consumer.accept(createPartialEvent("Processing data...", true, false), agentCard); + consumer.accept(createMessageEvent("Standard chat update message"), agentCard); + consumer.accept(createFinalEvent("Done"), agentCard); + }); + + List events = agent.runAsync(invocationContext).toList().blockingGet(); + + // The stream must not terminate early on MessageEvent, and run until createFinalEvent. + assertThat(events).hasSize(4); + assertText(events.get(0), "Processing data..."); + assertText(events.get(1), "Processing data..."); // aggregated (flushed) + assertText(events.get(2), "Standard chat update message"); // message event + assertText(events.get(3), "Done"); // terminal completed event + } + private RemoteA2AAgent.Builder getAgentBuilder() { return RemoteA2AAgent.builder().name("remote-agent").a2aClient(mockClient).agentCard(agentCard); }