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
18 changes: 7 additions & 11 deletions a2a/src/main/java/com/google/adk/a2a/agent/RemoteA2AAgent.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -227,8 +227,7 @@ protected Flowable<Event> 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<BiConsumer<ClientEvent, AgentCard>> consumers =
ImmutableList.of(handler::handleEvent);
a2aClient.sendMessage(originalMessage, consumers, handler::handleError, null);
Expand All @@ -249,7 +248,6 @@ private static class StreamHandler {
private final FlowableEmitter<Event> 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();
Expand All @@ -259,12 +257,10 @@ private static class StreamHandler {
FlowableEmitter<Event> emitter,
InvocationContext invocationContext,
String requestJson,
boolean streaming,
String agentName) {
this.emitter = emitter;
this.invocationContext = invocationContext;
this.requestJson = requestJson;
this.streaming = streaming;
this.agentName = agentName;
}

Expand Down Expand Up @@ -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<Part> eventParts(Event event) {
Expand Down
75 changes: 75 additions & 0 deletions a2a/src/test/java/com/google/adk/a2a/agent/RemoteA2AAgentTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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 =
Expand Down Expand Up @@ -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<Event> 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<Event> 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<Event> 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);
}
Expand Down
Loading