From b4e0d40f044d650f7fdde8a7b27f01e36ec61854 Mon Sep 17 00:00:00 2001 From: Kabir Khan Date: Tue, 4 Aug 2026 15:48:49 +0100 Subject: [PATCH] feat: add TaskStreamLifecycleHook for stream lifecycle management Add a CDI-discoverable hook that lets users observe task stream lifecycle events (subscribe, unsubscribe, event) and close all ChildQueues for a task on demand via StreamCloseHandle. Wired through InMemoryQueueManager, MainEventBusProcessor, and ReplicatedQueueManager. Includes a stream-lifecycle example with integration tests for all three transports (JSONRPC, gRPC, REST) and documentation. This fixes #990 Co-Authored-By: Claude Opus 4.6 (1M context) --- docs/content/dev/examples.md | 39 +++ docs/content/dev/server.md | 65 ++++ examples/stream-lifecycle/README.md | 135 +++++++++ examples/stream-lifecycle/client/pom.xml | 58 ++++ .../client/StreamLifecycleClient.java | 280 ++++++++++++++++++ examples/stream-lifecycle/pom.xml | 40 +++ examples/stream-lifecycle/server/pom.xml | 85 ++++++ .../server/AgentCardProducer.java | 53 ++++ .../server/AgentExecutorProducer.java | 59 ++++ .../server/CloseStreamsHook.java | 52 ++++ .../src/main/resources/application.properties | 4 + .../server/StreamLifecycleTest.java | 59 ++++ .../src/test/resources/application.properties | 3 + .../core/ReplicatedQueueManager.java | 17 +- pom.xml | 1 + .../DefaultTaskStreamLifecycleHook.java | 25 ++ .../sdk/server/events/EventQueue.java | 59 ++++ .../server/events/InMemoryQueueManager.java | 22 +- .../server/events/MainEventBusProcessor.java | 10 + .../sdk/server/events/StreamCloseHandle.java | 25 ++ .../events/TaskStreamLifecycleHook.java | 77 +++++ .../sdk/server/events/EventQueueTest.java | 100 +++++++ .../events/InMemoryQueueManagerTest.java | 97 ++++++ 23 files changed, 1362 insertions(+), 3 deletions(-) create mode 100644 examples/stream-lifecycle/README.md create mode 100644 examples/stream-lifecycle/client/pom.xml create mode 100644 examples/stream-lifecycle/client/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/client/StreamLifecycleClient.java create mode 100644 examples/stream-lifecycle/pom.xml create mode 100644 examples/stream-lifecycle/server/pom.xml create mode 100644 examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/AgentCardProducer.java create mode 100644 examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/AgentExecutorProducer.java create mode 100644 examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/CloseStreamsHook.java create mode 100644 examples/stream-lifecycle/server/src/main/resources/application.properties create mode 100644 examples/stream-lifecycle/server/src/test/java/org/a2aproject/sdk/examples/streamlifecycle/server/StreamLifecycleTest.java create mode 100644 examples/stream-lifecycle/server/src/test/resources/application.properties create mode 100644 server-common/src/main/java/org/a2aproject/sdk/server/events/DefaultTaskStreamLifecycleHook.java create mode 100644 server-common/src/main/java/org/a2aproject/sdk/server/events/StreamCloseHandle.java create mode 100644 server-common/src/main/java/org/a2aproject/sdk/server/events/TaskStreamLifecycleHook.java diff --git a/docs/content/dev/examples.md b/docs/content/dev/examples.md index 9d5b0fb40..949403633 100644 --- a/docs/content/dev/examples.md +++ b/docs/content/dev/examples.md @@ -144,6 +144,45 @@ The client expects an OpenTelemetry collector on port 5317. The easiest way is t For more information, see the [OpenTelemetry extras module](extras/opentelemetry). +## Stream Lifecycle Hook + +This example demonstrates the `TaskStreamLifecycleHook` — a CDI-discoverable hook that observes stream lifecycle events and can close all active streams on demand. + +The server implements a hook that closes all subscriber streams when 3 clients connect to the same task. The client creates 3 subscribers sequentially, sending messages while the first two are active. When the third subscriber connects, the hook fires and all streams close gracefully. + +### Start the Server + +```bash +cd examples/stream-lifecycle/server +mvn quarkus:dev +``` + +### Run the Client + +```bash +cd examples/stream-lifecycle/client +mvn exec:java +``` + +The client logs each event received by each subscriber, showing events flowing to subscribers 1 and 2 before the hook triggers. + +#### Transport Protocol Selection + +```bash +# JSON-RPC (default) +mvn exec:java + +# gRPC +mvn exec:java -Dquarkus.agentcard.protocol=GRPC + +# HTTP+JSON/REST +mvn exec:java -Dquarkus.agentcard.protocol=HTTP+JSON +``` + +Select the same protocol on both server and client. + +For implementation details, see the [Stream Lifecycle Hook section](server#6-stream-lifecycle-hook-optional) in the Server Guide. + ## More Examples - [a2a-samples repository](https://github.com/a2aproject/a2a-samples/tree/main/samples/java/agents) — Additional agent examples in Java and other languages diff --git a/docs/content/dev/server.md b/docs/content/dev/server.md index 39fe727dd..de9a0c6a7 100644 --- a/docs/content/dev/server.md +++ b/docs/content/dev/server.md @@ -158,6 +158,71 @@ See [Configuration](configuration) for all config properties and tuning. See [Task Authorization](authorization) for per-user access control. +## 6. Stream Lifecycle Hook (Optional) + +The `TaskStreamLifecycleHook` lets you observe and control streaming connections for a task. You are notified when clients subscribe, unsubscribe, or when events are distributed, and you can close all active streams on demand via the `StreamCloseHandle`. + +### Implementing a Hook + +Create a CDI bean that implements `TaskStreamLifecycleHook` and overrides the default no-op: + +```java +@ApplicationScoped +@Alternative +@Priority(1) +public class MyStreamHook implements TaskStreamLifecycleHook { + + @Override + public void onSubscribe(String taskId, StreamCloseHandle handle) { + // Called when a client subscribes to a task's event stream + } + + @Override + public void onUnsubscribe(String taskId, StreamCloseHandle handle) { + // Called when a client disconnects + } + + @Override + public void onEvent(String taskId, Event event, StreamCloseHandle handle) { + // Called after an event is persisted and distributed to all subscribers + } +} +``` + +### StreamCloseHandle + +Each callback receives a `StreamCloseHandle` with two methods: + +- **`closeStreams()`** — Gracefully closes all active subscriber streams for the task. The agent executor continues running and the MainQueue stays alive (for non-finalized tasks), so new clients can resubscribe. +- **`getActiveSubscriberCount()`** — Returns the number of currently connected subscribers. + +### Example: Close Streams at a Subscriber Threshold + +```java +@ApplicationScoped +@Alternative +@Priority(1) +public class CloseStreamsHook implements TaskStreamLifecycleHook { + + private static final int MAX_SUBSCRIBERS = 3; + + @Override + public void onSubscribe(String taskId, StreamCloseHandle handle) { + if (handle.getActiveSubscriberCount() >= MAX_SUBSCRIBERS) { + handle.closeStreams(); + } + } + + @Override + public void onUnsubscribe(String taskId, StreamCloseHandle handle) { } + + @Override + public void onEvent(String taskId, Event event, StreamCloseHandle handle) { } +} +``` + +See the [`examples/stream-lifecycle`](https://github.com/a2aproject/a2a-java/tree/main/examples/stream-lifecycle) directory for a complete working example with server, client, and integration tests for all three transports. + ## Backward Compatibility with v0.3 See [Backward Compatibility](compatibility) for multi-version modules, version routing, and v0.3 client support. diff --git a/examples/stream-lifecycle/README.md b/examples/stream-lifecycle/README.md new file mode 100644 index 000000000..72db235c0 --- /dev/null +++ b/examples/stream-lifecycle/README.md @@ -0,0 +1,135 @@ +# Stream Lifecycle Hook Example + +This example demonstrates `TaskStreamLifecycleHook`, a CDI-discoverable hook that lets you observe and control streaming connections for a task. The server closes all active streams when 3 subscribers connect to the same task. + +For example, in order to save resources associated with streaming connections, you might want to close streams: + +* If no events are received for a Task in a given timeframe. This can help if you have a lot of Tasks taking a very long time to complete +* Stopping rogue client opening too many subscriptions to the same Task + +## Prerequisites + +- Java 17 or higher +- Maven + +## What It Does + +**Server** — An agent sends 20 progress messages (one every 500ms). A `CloseStreamsHook` monitors subscriber count and calls `StreamCloseHandle.closeStreams()` when 3 subscribers are connected, gracefully closing all streams. + +**Client** — Connects 3 subscribers sequentially: +1. Subscriber 1 sends a message (creates the task, starts streaming) +2. Subscriber 2 subscribes to the same task after 1.5 seconds +3. Both subscribers receive progress messages for 2 seconds +4. Subscriber 3 subscribes — the hook fires and closes all streams +5. All 3 subscribers see their streams end gracefully + +## Run the Example + +### 1. Build the SDK + +From the repository root: + +```bash +mvn clean install -DskipTests +``` + +### 2. Start the Server + +```bash +cd examples/stream-lifecycle/server +mvn quarkus:dev +``` + +### 3. Run the Client + +In a separate terminal: + +```bash +cd examples/stream-lifecycle/client +mvn exec:java +``` + +### Expected Output (Client) + +The exact event interleaving depends on timing, but you should see something like: + +``` +Resolved agent card: Stream Lifecycle Demo Agent +[Sub-1] Sending message to create task... +[Sub-1] StatusUpdate — TASK_STATE_WORKING + +=== Task created: === + +[Sub-1] ArtifactEvent — Progress update 1/20 +[Sub-1] ArtifactEvent — Progress update 2/20 +[Sub-1] ArtifactEvent — Progress update 3/20 +[Sub-2] Subscribing to task ... + +=== Subscribers 1 and 2 are active — receiving events... === + +[Sub-2] TaskEvent — state: TASK_STATE_WORKING, id: +[Sub-1] ArtifactEvent — Progress update 4/20 +[Sub-2] ArtifactEvent — Progress update 4/20 +[Sub-1] ArtifactEvent — Progress update 5/20 +[Sub-2] ArtifactEvent — Progress update 5/20 +... +[Sub-3] Subscribing to task (will trigger stream close)... +[Sub-1] Stream closed. +[Sub-2] Stream closed. +[Sub-3] Stream closed. + +=== All streams closed. === +``` + +### Expected Output (Server) + +``` +[HOOK] Subscriber added for task . Active subscribers: 1 +[AGENT] Starting execution for task +[AGENT] Sending: Progress update 1/20 +[HOOK] Event distributed for task : Message (subscribers: 1) +... +[HOOK] Subscriber added for task . Active subscribers: 2 +... +[HOOK] Subscriber added for task . Active subscribers: 3 +[HOOK] Subscriber count reached 3 for task — closing all streams +[HOOK] Subscriber removed for task . Active subscribers: 2 +[HOOK] Subscriber removed for task . Active subscribers: 1 +[HOOK] Subscriber removed for task . Active subscribers: 0 +``` + +## Transport Protocol Selection + +Set `quarkus.agentcard.protocol` on both server and client (must match). Available values: + +| Value | Transport | +|-------|-----------| +| `JSONRPC` | JSON-RPC 2.0 (default) | +| `GRPC` | gRPC | +| `HTTP+JSON` | HTTP+JSON/REST | + +```bash +# Server — gRPC example +mvn quarkus:dev -Dquarkus.agentcard.protocol=GRPC + +# Client — must use the same value +mvn exec:java -Dquarkus.agentcard.protocol=GRPC +``` + +## Key Files + +| File | Description | +|------|-------------| +| `server/.../CloseStreamsHook.java` | `TaskStreamLifecycleHook` implementation — closes streams at 3 subscribers | +| `server/.../AgentExecutorProducer.java` | Agent that sends 20 progress messages over 10 seconds | +| `server/.../AgentCardProducer.java` | Agent card with streaming enabled | +| `client/.../StreamLifecycleClient.java` | Client that creates 3 subscribers and logs events | + +## Integration Tests + +The server module includes `@QuarkusTest` integration tests that verify the hook behavior across all three transports: + +```bash +cd examples/stream-lifecycle/server +mvn test +``` diff --git a/examples/stream-lifecycle/client/pom.xml b/examples/stream-lifecycle/client/pom.xml new file mode 100644 index 000000000..244d304f3 --- /dev/null +++ b/examples/stream-lifecycle/client/pom.xml @@ -0,0 +1,58 @@ + + + 4.0.0 + + + org.a2aproject.sdk + a2a-java-sdk-examples-stream-lifecycle-parent + 1.1.1.Final-SNAPSHOT + + + a2a-java-sdk-examples-stream-lifecycle-client + + Java SDK A2A Examples - Stream Lifecycle Client + Client demonstrating TaskStreamLifecycleHook with multiple subscribers + + + + org.a2aproject.sdk + a2a-java-sdk-client + + + org.a2aproject.sdk + a2a-java-sdk-jsonrpc-common + + + org.a2aproject.sdk + a2a-java-sdk-client-transport-grpc + + + io.grpc + grpc-netty + + + org.a2aproject.sdk + a2a-java-sdk-client-transport-rest + + + org.junit.jupiter + junit-jupiter-api + compile + + + + + + + org.codehaus.mojo + exec-maven-plugin + 3.6.3 + + org.a2aproject.sdk.examples.streamlifecycle.client.StreamLifecycleClient + + + + + diff --git a/examples/stream-lifecycle/client/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/client/StreamLifecycleClient.java b/examples/stream-lifecycle/client/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/client/StreamLifecycleClient.java new file mode 100644 index 000000000..b1d7df945 --- /dev/null +++ b/examples/stream-lifecycle/client/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/client/StreamLifecycleClient.java @@ -0,0 +1,280 @@ +package org.a2aproject.sdk.examples.streamlifecycle.client; + +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.BiConsumer; +import java.util.function.Function; + +import org.a2aproject.sdk.A2A; +import org.a2aproject.sdk.client.Client; +import org.a2aproject.sdk.client.ClientBuilder; +import org.a2aproject.sdk.client.ClientEvent; +import org.a2aproject.sdk.client.MessageEvent; +import org.a2aproject.sdk.client.TaskEvent; +import org.a2aproject.sdk.client.TaskUpdateEvent; +import org.a2aproject.sdk.client.config.ClientConfig; +import org.a2aproject.sdk.client.http.A2ACardResolver; +import org.a2aproject.sdk.client.transport.grpc.GrpcTransport; +import org.a2aproject.sdk.client.transport.grpc.GrpcTransportConfigBuilder; +import org.a2aproject.sdk.client.transport.jsonrpc.JSONRPCTransport; +import org.a2aproject.sdk.client.transport.jsonrpc.JSONRPCTransportConfig; +import org.a2aproject.sdk.client.transport.rest.RestTransport; +import org.a2aproject.sdk.client.transport.rest.RestTransportConfig; +import org.a2aproject.sdk.client.transport.spi.ClientTransportConfig; +import org.a2aproject.sdk.spec.AgentCard; +import org.a2aproject.sdk.spec.Message; +import org.a2aproject.sdk.spec.Part; +import org.a2aproject.sdk.spec.Task; +import org.a2aproject.sdk.spec.TaskArtifactUpdateEvent; +import org.a2aproject.sdk.spec.TaskIdParams; +import org.a2aproject.sdk.spec.TextPart; + +import io.grpc.Channel; +import io.grpc.ManagedChannel; +import io.grpc.ManagedChannelBuilder; + +/** + * Demonstrates the TaskStreamLifecycleHook by connecting 3 subscribers to a task. + *

+ * Flow: + *

    + *
  1. Subscriber 1: sends a message (creates the task, starts streaming)
  2. + *
  3. Subscriber 2: subscribes to the same task
  4. + *
  5. Both subscribers receive progress messages from the agent
  6. + *
  7. Subscriber 3: subscribes — the server hook detects 3 subscribers and closes all streams
  8. + *
  9. All subscribers see their streams end gracefully
  10. + *
+ *

+ * Start the server first ({@code mvn quarkus:dev} in the server module), then run this client. + *

+ * This class also serves as the integration test body — {@code runAndVerify(AgentCard, String)} + * is called by the server module's {@code @QuarkusTest} for each transport protocol. + */ +public class StreamLifecycleClient { + + private static final String SERVER_URL = "http://localhost:9999"; + + private static final List grpcChannels = new CopyOnWriteArrayList<>(); + + public static void main(String[] args) throws Exception { + try { + AgentCard agentCard = A2ACardResolver.builder().baseUrl(SERVER_URL).build().getAgentCard(); + System.out.println("Resolved agent card: " + agentCard.name()); + String protocol = System.getProperty("quarkus.agentcard.protocol", "JSONRPC"); + runAndVerify(agentCard, protocol); + } finally { + shutdownGrpcChannels(); + } + } + + /** + * Runs the 3-subscriber scenario and asserts correctness. + * Called by {@code main()} for manual runs and by the server's {@code @QuarkusTest} for each protocol. + */ + public static void runAndVerify(AgentCard agentCard, String protocol) throws Exception { + AtomicReference taskIdRef = new AtomicReference<>(); + CountDownLatch taskIdReady = new CountDownLatch(1); + + CountDownLatch sub1StreamDone = new CountDownLatch(1); + CountDownLatch sub2StreamDone = new CountDownLatch(1); + CountDownLatch sub3StreamDone = new CountDownLatch(1); + + CopyOnWriteArrayList sub1Events = new CopyOnWriteArrayList<>(); + CopyOnWriteArrayList sub2Events = new CopyOnWriteArrayList<>(); + + // --- Subscriber 1: Send a message (creates the task) --- + Thread sub1Thread = new Thread(() -> { + try { + Client client = createStreamingClient(agentCard, "Sub-1", (event, card) -> { + sub1Events.add(event); + logEvent("Sub-1", event); + if (taskIdRef.get() == null) { + String id = extractTaskId(event); + if (id != null) { + taskIdRef.set(id); + taskIdReady.countDown(); + } + } + }, sub1StreamDone, protocol); + + System.out.println("[Sub-1] Sending message to create task..."); + client.sendMessage(A2A.toUserMessage("Start streaming demo")); + } catch (Exception e) { + System.out.println("[Sub-1] Error: " + e.getMessage()); + } finally { + sub1StreamDone.countDown(); + } + }, "subscriber-1"); + sub1Thread.start(); + + assertTrue(taskIdReady.await(10, TimeUnit.SECONDS), "Task should be created"); + String taskId = taskIdRef.get(); + assertNotNull(taskId); + System.out.println("\n=== Task created: " + taskId + " ===\n"); + + // Give some time for messages to flow to subscriber 1 + Thread.sleep(1500); + + // --- Subscriber 2: Subscribe to the existing task --- + Thread sub2Thread = new Thread(() -> { + try { + Client client = createStreamingClient(agentCard, "Sub-2", (event, card) -> { + sub2Events.add(event); + logEvent("Sub-2", event); + }, sub2StreamDone, protocol); + + System.out.println("[Sub-2] Subscribing to task " + taskId + "..."); + client.subscribeToTask(new TaskIdParams(taskId)); + } catch (Exception e) { + System.out.println("[Sub-2] Error: " + e.getMessage()); + } finally { + sub2StreamDone.countDown(); + } + }, "subscriber-2"); + sub2Thread.start(); + + System.out.println("\n=== Subscribers 1 and 2 are active — receiving events... ===\n"); + Thread.sleep(2000); + + // --- Subscriber 3: Triggers the hook (closes all streams) --- + Thread sub3Thread = new Thread(() -> { + try { + Client client = createStreamingClient(agentCard, "Sub-3", (event, card) -> { + logEvent("Sub-3", event); + }, sub3StreamDone, protocol); + + System.out.println("[Sub-3] Subscribing to task " + taskId + " (will trigger stream close)..."); + client.subscribeToTask(new TaskIdParams(taskId)); + } catch (Exception e) { + System.out.println("[Sub-3] Error: " + e.getMessage()); + } finally { + sub3StreamDone.countDown(); + } + }, "subscriber-3"); + sub3Thread.start(); + + // Wait for all streams to close + assertTrue(sub1StreamDone.await(30, TimeUnit.SECONDS), "Subscriber 1 stream should close"); + assertTrue(sub2StreamDone.await(30, TimeUnit.SECONDS), "Subscriber 2 stream should close"); + assertTrue(sub3StreamDone.await(30, TimeUnit.SECONDS), "Subscriber 3 stream should close"); + + System.out.println("\n=== All streams closed. ==="); + + // Verify events were received + assertTrue(sub1Events.size() >= 2, + "Subscriber 1 should have received at least a status update and one artifact, got: " + sub1Events.size()); + assertTrue(sub2Events.size() >= 1, + "Subscriber 2 should have received at least a task snapshot, got: " + sub2Events.size()); + } + + public static void shutdownGrpcChannels() { + for (ManagedChannel ch : grpcChannels) { + ch.shutdownNow(); + try { + ch.awaitTermination(5, TimeUnit.SECONDS); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + grpcChannels.clear(); + } + + private static Client createStreamingClient(AgentCard agentCard, String name, + BiConsumer consumer, + CountDownLatch streamDone, + String protocol) throws Exception { + ClientBuilder builder = Client.builder(agentCard) + .addConsumer(consumer) + .streamingErrorHandler(e -> { + System.out.println("[" + name + "] Stream closed."); + streamDone.countDown(); + }) + .clientConfig(ClientConfig.builder().setStreaming(true).build()); + configureTransport(builder, protocol); + return builder.build(); + } + + private static void configureTransport(ClientBuilder clientBuilder, String protocol) { + ClientTransportConfig transportConfig; + switch (protocol) { + case "GRPC": + Function channelFactory = url -> { + String target = url.replaceAll("^https?://", ""); + ManagedChannel channel = ManagedChannelBuilder.forTarget(target).usePlaintext().build(); + grpcChannels.add(channel); + return channel; + }; + transportConfig = new GrpcTransportConfigBuilder().channelFactory(channelFactory).build(); + clientBuilder.withTransport(GrpcTransport.class, transportConfig); + break; + case "HTTP+JSON": + transportConfig = new RestTransportConfig(); + clientBuilder.withTransport(RestTransport.class, transportConfig); + break; + case "JSONRPC": + default: + transportConfig = new JSONRPCTransportConfig(); + clientBuilder.withTransport(JSONRPCTransport.class, transportConfig); + break; + } + } + + public static void logEvent(String subscriber, ClientEvent event) { + if (event instanceof TaskEvent te) { + System.out.printf("[%s] TaskEvent — state: %s, id: %s%n", + subscriber, te.getTask().status().state(), te.getTask().id()); + } else if (event instanceof TaskUpdateEvent tue) { + if (tue.getUpdateEvent() instanceof TaskArtifactUpdateEvent artifact) { + String text = extractText(artifact); + System.out.printf("[%s] ArtifactEvent — %s%n", subscriber, text); + } else { + System.out.printf("[%s] StatusUpdate — %s%n", subscriber, tue.getTask().status().state()); + } + } else if (event instanceof MessageEvent me) { + String text = extractText(me.getMessage()); + System.out.printf("[%s] MessageEvent — %s%n", subscriber, text); + } + } + + static String extractTaskId(ClientEvent event) { + if (event instanceof TaskEvent te) { + return te.getTask().id(); + } else if (event instanceof TaskUpdateEvent tue) { + Task task = tue.getTask(); + return task != null ? task.id() : null; + } + return null; + } + + private static String extractText(TaskArtifactUpdateEvent artifact) { + if (artifact.artifact() == null || artifact.artifact().parts() == null) { + return "(no parts)"; + } + StringBuilder sb = new StringBuilder(); + for (Part part : artifact.artifact().parts()) { + if (part instanceof TextPart tp) { + sb.append(tp.text()); + } + } + return sb.toString(); + } + + private static String extractText(Message message) { + if (message.parts() == null) { + return "(no parts)"; + } + StringBuilder sb = new StringBuilder(); + for (Part part : message.parts()) { + if (part instanceof TextPart tp) { + sb.append(tp.text()); + } + } + return sb.toString(); + } +} diff --git a/examples/stream-lifecycle/pom.xml b/examples/stream-lifecycle/pom.xml new file mode 100644 index 000000000..10d180304 --- /dev/null +++ b/examples/stream-lifecycle/pom.xml @@ -0,0 +1,40 @@ + + + 4.0.0 + + + org.a2aproject.sdk + a2a-java-sdk-parent + 1.1.1.Final-SNAPSHOT + ../../pom.xml + + + a2a-java-sdk-examples-stream-lifecycle-parent + pom + + Java SDK A2A Examples: Stream Lifecycle + Example demonstrating TaskStreamLifecycleHook for managing streaming connections + + + + + io.quarkus + quarkus-bom + ${quarkus.platform.version} + pom + import + + + org.a2aproject.sdk + a2a-java-sdk-client + + + + + + client + server + + diff --git a/examples/stream-lifecycle/server/pom.xml b/examples/stream-lifecycle/server/pom.xml new file mode 100644 index 000000000..b26ed1f3f --- /dev/null +++ b/examples/stream-lifecycle/server/pom.xml @@ -0,0 +1,85 @@ + + + 4.0.0 + + + org.a2aproject.sdk + a2a-java-sdk-examples-stream-lifecycle-parent + 1.1.1.Final-SNAPSHOT + + + a2a-java-sdk-examples-stream-lifecycle-server + + Java SDK A2A Examples - Stream Lifecycle Server + Server demonstrating TaskStreamLifecycleHook for closing streams on demand + + + + org.a2aproject.sdk + a2a-java-sdk-client + + + org.a2aproject.sdk + a2a-java-sdk-reference-jsonrpc + + + io.quarkus + quarkus-resteasy + provided + + + org.a2aproject.sdk + a2a-java-sdk-reference-grpc + + + org.a2aproject.sdk + a2a-java-sdk-reference-rest + + + jakarta.enterprise + jakarta.enterprise.cdi-api + provided + + + jakarta.ws.rs + jakarta.ws.rs-api + + + + + io.quarkus + quarkus-junit5 + test + + + org.a2aproject.sdk + a2a-java-sdk-examples-stream-lifecycle-client + ${project.version} + test + + + + + + + io.quarkus + quarkus-maven-plugin + true + + + + build + generate-code + generate-code-tests + + + + + --add-opens=java.base/java.lang=ALL-UNNAMED + + + + + diff --git a/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/AgentCardProducer.java b/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/AgentCardProducer.java new file mode 100644 index 000000000..807a26fa8 --- /dev/null +++ b/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/AgentCardProducer.java @@ -0,0 +1,53 @@ +package org.a2aproject.sdk.examples.streamlifecycle.server; + +import java.util.Collections; +import java.util.List; + +import jakarta.enterprise.context.ApplicationScoped; +import jakarta.enterprise.inject.Produces; + +import org.a2aproject.sdk.server.PublicAgentCard; +import org.a2aproject.sdk.spec.AgentCapabilities; +import org.a2aproject.sdk.spec.AgentCard; +import org.a2aproject.sdk.spec.AgentInterface; +import org.a2aproject.sdk.spec.AgentSkill; +import org.a2aproject.sdk.spec.TransportProtocol; +import org.eclipse.microprofile.config.inject.ConfigProperty; + +@ApplicationScoped +public class AgentCardProducer { + + @ConfigProperty(name = "quarkus.agentcard.protocol", defaultValue = "JSONRPC") + String protocolStr; + + @Produces + @PublicAgentCard + public AgentCard agentCard() { + return AgentCard.builder() + .name("Stream Lifecycle Demo Agent") + .description("Demonstrates TaskStreamLifecycleHook — closes all streams when 3 subscribers connect") + .supportedInterfaces(Collections.singletonList(getAgentInterface())) + .version("1.0.0") + .capabilities(AgentCapabilities.builder() + .streaming(true) + .build()) + .defaultInputModes(Collections.singletonList("text")) + .defaultOutputModes(Collections.singletonList("text")) + .skills(Collections.singletonList(AgentSkill.builder() + .id("stream_lifecycle") + .name("Stream lifecycle demo") + .description("Sends progress messages while subscribers connect and disconnect") + .tags(List.of("streaming", "lifecycle")) + .build())) + .build(); + } + + private AgentInterface getAgentInterface() { + TransportProtocol protocol = TransportProtocol.fromString(protocolStr); + String url = switch (protocol) { + case GRPC -> "localhost:9000"; + case JSONRPC, HTTP_JSON -> "http://localhost:9999"; + }; + return new AgentInterface(protocol.asString(), url); + } +} diff --git a/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/AgentExecutorProducer.java b/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/AgentExecutorProducer.java new file mode 100644 index 000000000..e150aa57d --- /dev/null +++ b/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/AgentExecutorProducer.java @@ -0,0 +1,59 @@ +package org.a2aproject.sdk.examples.streamlifecycle.server; + +import jakarta.enterprise.context.ApplicationScoped; +import jakarta.enterprise.inject.Produces; + +import org.a2aproject.sdk.server.agentexecution.AgentExecutor; +import org.a2aproject.sdk.server.agentexecution.RequestContext; +import org.a2aproject.sdk.server.tasks.AgentEmitter; +import org.a2aproject.sdk.spec.A2AError; +import org.a2aproject.sdk.spec.TextPart; +import org.a2aproject.sdk.spec.UnsupportedOperationError; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * Agent executor that sends a series of messages over time, giving clients time to subscribe. + *

+ * The agent sends 20 progress messages, one every 500ms. This gives the demo client enough + * time to establish multiple subscriptions and observe events flowing to each subscriber + * before the {@link CloseStreamsHook} closes all streams. + *

+ */ +@ApplicationScoped +public class AgentExecutorProducer { + + private static final Logger LOG = LoggerFactory.getLogger(AgentExecutorProducer.class); + + @Produces + public AgentExecutor agentExecutor() { + return new AgentExecutor() { + @Override + public void execute(RequestContext context, AgentEmitter emitter) throws A2AError { + LOG.info("[AGENT] Starting execution for task {}", context.getTaskId()); + emitter.startWork(); + + for (int i = 1; i <= 20; i++) { + try { + Thread.sleep(500); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + LOG.info("[AGENT] Interrupted at message {}", i); + break; + } + String text = "Progress update " + i + "/20"; + LOG.info("[AGENT] Sending: {}", text); + emitter.addArtifact(java.util.List.of(new TextPart(text))); + } + + LOG.info("[AGENT] Completing task {}", context.getTaskId()); + emitter.complete(); + } + + @Override + public void cancel(RequestContext context, AgentEmitter emitter) throws A2AError { + throw new UnsupportedOperationError(); + } + }; + } +} diff --git a/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/CloseStreamsHook.java b/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/CloseStreamsHook.java new file mode 100644 index 000000000..2c3b5354c --- /dev/null +++ b/examples/stream-lifecycle/server/src/main/java/org/a2aproject/sdk/examples/streamlifecycle/server/CloseStreamsHook.java @@ -0,0 +1,52 @@ +package org.a2aproject.sdk.examples.streamlifecycle.server; + +import jakarta.annotation.Priority; +import jakarta.enterprise.context.ApplicationScoped; +import jakarta.enterprise.inject.Alternative; + +import org.a2aproject.sdk.server.events.StreamCloseHandle; +import org.a2aproject.sdk.server.events.TaskStreamLifecycleHook; +import org.a2aproject.sdk.spec.Event; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * Demonstrates closing all streams for a task when the subscriber count reaches a threshold. + *

+ * When 3 subscribers are connected to the same task, this hook calls + * {@link StreamCloseHandle#closeStreams()} to gracefully close all active streams. + * The agent executor continues running, but disconnected clients can resubscribe later. + *

+ */ +@ApplicationScoped +@Alternative +@Priority(1) +public class CloseStreamsHook implements TaskStreamLifecycleHook { + + private static final Logger LOG = LoggerFactory.getLogger(CloseStreamsHook.class); + private static final int MAX_SUBSCRIBERS = 3; + + @Override + public void onSubscribe(String taskId, StreamCloseHandle handle) { + int count = handle.getActiveSubscriberCount(); + LOG.info("[HOOK] Subscriber added for task {}. Active subscribers: {}", taskId, count); + + if (count >= MAX_SUBSCRIBERS) { + LOG.info("[HOOK] Subscriber count reached {} for task {} — closing all streams", + MAX_SUBSCRIBERS, taskId); + handle.closeStreams(); + } + } + + @Override + public void onUnsubscribe(String taskId, StreamCloseHandle handle) { + LOG.info("[HOOK] Subscriber removed for task {}. Active subscribers: {}", + taskId, handle.getActiveSubscriberCount()); + } + + @Override + public void onEvent(String taskId, Event event, StreamCloseHandle handle) { + LOG.info("[HOOK] Event distributed for task {}: {} (subscribers: {})", + taskId, event.getClass().getSimpleName(), handle.getActiveSubscriberCount()); + } +} diff --git a/examples/stream-lifecycle/server/src/main/resources/application.properties b/examples/stream-lifecycle/server/src/main/resources/application.properties new file mode 100644 index 000000000..dfea3ba12 --- /dev/null +++ b/examples/stream-lifecycle/server/src/main/resources/application.properties @@ -0,0 +1,4 @@ +%dev.quarkus.http.port=9999 + +# Protocol can be JSONRPC, GRPC, or HTTP+JSON +quarkus.agentcard.protocol=JSONRPC diff --git a/examples/stream-lifecycle/server/src/test/java/org/a2aproject/sdk/examples/streamlifecycle/server/StreamLifecycleTest.java b/examples/stream-lifecycle/server/src/test/java/org/a2aproject/sdk/examples/streamlifecycle/server/StreamLifecycleTest.java new file mode 100644 index 000000000..07258a975 --- /dev/null +++ b/examples/stream-lifecycle/server/src/test/java/org/a2aproject/sdk/examples/streamlifecycle/server/StreamLifecycleTest.java @@ -0,0 +1,59 @@ +package org.a2aproject.sdk.examples.streamlifecycle.server; + +import java.util.List; +import java.util.concurrent.TimeUnit; + +import org.a2aproject.sdk.examples.streamlifecycle.client.StreamLifecycleClient; +import org.a2aproject.sdk.spec.AgentCapabilities; +import org.a2aproject.sdk.spec.AgentCard; +import org.a2aproject.sdk.spec.AgentInterface; +import org.a2aproject.sdk.spec.TransportProtocol; + +import io.quarkus.test.junit.QuarkusTest; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; + +@QuarkusTest +class StreamLifecycleTest { + + private static final String HTTP_URL = "http://localhost:8081"; + private static final String GRPC_URL = "localhost:8081"; + + @AfterEach + void cleanupGrpcChannels() { + StreamLifecycleClient.shutdownGrpcChannels(); + } + + @Test + @Timeout(value = 60, unit = TimeUnit.SECONDS) + void testStreamLifecycleWithJsonRpc() throws Exception { + StreamLifecycleClient.runAndVerify(buildCard(TransportProtocol.JSONRPC, HTTP_URL), "JSONRPC"); + } + + @Test + @Timeout(value = 60, unit = TimeUnit.SECONDS) + void testStreamLifecycleWithRest() throws Exception { + StreamLifecycleClient.runAndVerify(buildCard(TransportProtocol.HTTP_JSON, HTTP_URL), "HTTP+JSON"); + } + + @Test + @Timeout(value = 60, unit = TimeUnit.SECONDS) + void testStreamLifecycleWithGrpc() throws Exception { + StreamLifecycleClient.runAndVerify(buildCard(TransportProtocol.GRPC, GRPC_URL), "GRPC"); + } + + private AgentCard buildCard(TransportProtocol protocol, String url) { + return AgentCard.builder() + .name("test") + .description("test") + .version("1.0.0") + .defaultInputModes(List.of("text")) + .defaultOutputModes(List.of("text")) + .skills(List.of()) + .supportedInterfaces(List.of( + new AgentInterface(protocol.asString(), url))) + .capabilities(AgentCapabilities.builder().streaming(true).build()) + .build(); + } +} diff --git a/examples/stream-lifecycle/server/src/test/resources/application.properties b/examples/stream-lifecycle/server/src/test/resources/application.properties new file mode 100644 index 000000000..8284b7631 --- /dev/null +++ b/examples/stream-lifecycle/server/src/test/resources/application.properties @@ -0,0 +1,3 @@ +quarkus.http.port=8081 +quarkus.grpc.server.use-separate-server=false +quarkus.agentcard.protocol=JSONRPC diff --git a/extras/queue-manager-replicated/core/src/main/java/org/a2aproject/sdk/extras/queuemanager/replicated/core/ReplicatedQueueManager.java b/extras/queue-manager-replicated/core/src/main/java/org/a2aproject/sdk/extras/queuemanager/replicated/core/ReplicatedQueueManager.java index 8ab19410d..1f1565563 100644 --- a/extras/queue-manager-replicated/core/src/main/java/org/a2aproject/sdk/extras/queuemanager/replicated/core/ReplicatedQueueManager.java +++ b/extras/queue-manager-replicated/core/src/main/java/org/a2aproject/sdk/extras/queuemanager/replicated/core/ReplicatedQueueManager.java @@ -8,6 +8,7 @@ import jakarta.inject.Inject; import org.a2aproject.sdk.extras.common.events.TaskFinalizedEvent; +import org.a2aproject.sdk.server.events.DefaultTaskStreamLifecycleHook; import org.a2aproject.sdk.server.events.EventEnqueueHook; import org.a2aproject.sdk.server.events.EventQueue; import org.a2aproject.sdk.server.events.EventQueueFactory; @@ -15,6 +16,7 @@ import org.a2aproject.sdk.server.events.InMemoryQueueManager; import org.a2aproject.sdk.server.events.MainEventBus; import org.a2aproject.sdk.server.events.QueueManager; +import org.a2aproject.sdk.server.events.TaskStreamLifecycleHook; import org.a2aproject.sdk.server.tasks.TaskStateProvider; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -31,6 +33,7 @@ public class ReplicatedQueueManager implements QueueManager { private InMemoryQueueManager delegate; private ReplicationStrategy replicationStrategy; private TaskStateProvider taskStateProvider; + private TaskStreamLifecycleHook streamLifecycleHook; /** * No-args constructor for CDI proxy creation. @@ -43,15 +46,25 @@ protected ReplicatedQueueManager() { this.delegate = null; this.replicationStrategy = null; this.taskStateProvider = null; + this.streamLifecycleHook = null; } - @Inject public ReplicatedQueueManager(ReplicationStrategy replicationStrategy, TaskStateProvider taskStateProvider, MainEventBus mainEventBus) { + this(replicationStrategy, taskStateProvider, mainEventBus, new DefaultTaskStreamLifecycleHook()); + } + + @Inject + public ReplicatedQueueManager(ReplicationStrategy replicationStrategy, + TaskStateProvider taskStateProvider, + MainEventBus mainEventBus, + TaskStreamLifecycleHook streamLifecycleHook) { this.replicationStrategy = replicationStrategy; this.taskStateProvider = taskStateProvider; - this.delegate = new InMemoryQueueManager(new ReplicatingEventQueueFactory(), taskStateProvider, mainEventBus); + this.streamLifecycleHook = streamLifecycleHook; + this.delegate = new InMemoryQueueManager( + new ReplicatingEventQueueFactory(), taskStateProvider, mainEventBus, streamLifecycleHook); } diff --git a/pom.xml b/pom.xml index 2f6f9f242..e9deeccd5 100644 --- a/pom.xml +++ b/pom.xml @@ -575,6 +575,7 @@ client/transport/spi common examples/helloworld + examples/stream-lifecycle examples/cloud-deployment/server extras/common extras/opentelemetry diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/events/DefaultTaskStreamLifecycleHook.java b/server-common/src/main/java/org/a2aproject/sdk/server/events/DefaultTaskStreamLifecycleHook.java new file mode 100644 index 000000000..8501aec80 --- /dev/null +++ b/server-common/src/main/java/org/a2aproject/sdk/server/events/DefaultTaskStreamLifecycleHook.java @@ -0,0 +1,25 @@ +package org.a2aproject.sdk.server.events; + +import jakarta.enterprise.context.ApplicationScoped; + +import org.a2aproject.sdk.spec.Event; + +/** + * Default no-op implementation of {@link TaskStreamLifecycleHook}. + * Override with a CDI {@code @Alternative @Priority} bean to provide custom stream lifecycle behavior. + */ +@ApplicationScoped +public class DefaultTaskStreamLifecycleHook implements TaskStreamLifecycleHook { + + @Override + public void onSubscribe(String taskId, StreamCloseHandle handle) { + } + + @Override + public void onUnsubscribe(String taskId, StreamCloseHandle handle) { + } + + @Override + public void onEvent(String taskId, Event event, StreamCloseHandle handle) { + } +} diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/events/EventQueue.java b/server-common/src/main/java/org/a2aproject/sdk/server/events/EventQueue.java index 253a0c8ad..d9bc13eef 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/events/EventQueue.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/events/EventQueue.java @@ -400,6 +400,19 @@ static class MainQueue extends EventQueue { private final List onCloseCallbacks; private final @Nullable TaskStateProvider taskStateProvider; private final MainEventBus mainEventBus; + private volatile @Nullable TaskStreamLifecycleHook streamLifecycleHook; + + private final StreamCloseHandle streamCloseHandle = new StreamCloseHandle() { + @Override + public void closeStreams() { + closeChildren(); + } + + @Override + public int getActiveSubscriberCount() { + return getActiveChildCount(); + } + }; MainQueue(int queueSize, @Nullable EventEnqueueHook hook, @@ -422,6 +435,13 @@ static class MainQueue extends EventQueue { public EventQueue tap() { ChildQueue child = new ChildQueue(this); children.add(child); + if (streamLifecycleHook != null) { + try { + streamLifecycleHook.onSubscribe(taskId, getStreamCloseHandle()); + } catch (Exception e) { + LOGGER.error("Error in TaskStreamLifecycleHook.onSubscribe for task {}", taskId, e); + } + } return child; } @@ -566,6 +586,15 @@ public void signalQueuePollerStarted() { void childClosing(ChildQueue child, boolean immediate) { children.remove(child); // Remove the closing child + // Notify stream lifecycle hook of unsubscription + if (streamLifecycleHook != null) { + try { + streamLifecycleHook.onUnsubscribe(taskId, getStreamCloseHandle()); + } catch (Exception e) { + LOGGER.error("Error in TaskStreamLifecycleHook.onUnsubscribe for task {}", taskId, e); + } + } + // If there are still children, keep queue open if (!children.isEmpty()) { LOGGER.debug("MainQueue staying open: {} children remaining", children.size()); @@ -629,6 +658,36 @@ public int getActiveChildCount() { return children.size(); } + void setTaskStreamLifecycleHook(@Nullable TaskStreamLifecycleHook hook) { + this.streamLifecycleHook = hook; + } + + @Nullable + TaskStreamLifecycleHook getTaskStreamLifecycleHook() { + return streamLifecycleHook; + } + + /** + * Gracefully closes all current ChildQueues without closing this MainQueue. + * Events already in ChildQueue deques will be drained before the stream terminates. + * New subscriptions via {@link #tap()} remain possible after this call. + *

+ * Note: if the task is finalized, the last {@code childClosing()} call will close the + * MainQueue itself (existing behavior). For non-finalized tasks, the MainQueue stays open. + *

+ */ + void closeChildren() { + List snapshot = List.copyOf(children); + LOGGER.debug("MainQueue[{}]: Closing {} children gracefully", taskId, snapshot.size()); + for (ChildQueue child : snapshot) { + child.close(false); + } + } + + StreamCloseHandle getStreamCloseHandle() { + return streamCloseHandle; + } + @Override protected void doClose(boolean immediate) { // Invoke all callbacks BEFORE closing, so they can still enqueue events diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/events/InMemoryQueueManager.java b/server-common/src/main/java/org/a2aproject/sdk/server/events/InMemoryQueueManager.java index 4ed55a0a8..5a87b65e8 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/events/InMemoryQueueManager.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/events/InMemoryQueueManager.java @@ -21,6 +21,7 @@ public class InMemoryQueueManager implements QueueManager { // final, is not proxyable in all runtimes private EventQueueFactory factory; private TaskStateProvider taskStateProvider; + private TaskStreamLifecycleHook streamLifecycleHook; /** * No-args constructor for CDI proxy creation. @@ -32,22 +33,35 @@ protected InMemoryQueueManager() { // For CDI proxy creation this.factory = null; this.taskStateProvider = null; + this.streamLifecycleHook = null; } MainEventBus mainEventBus; - @Inject public InMemoryQueueManager(TaskStateProvider taskStateProvider, MainEventBus mainEventBus) { + this(taskStateProvider, mainEventBus, new DefaultTaskStreamLifecycleHook()); + } + + @Inject + public InMemoryQueueManager(TaskStateProvider taskStateProvider, MainEventBus mainEventBus, + TaskStreamLifecycleHook streamLifecycleHook) { this.mainEventBus = mainEventBus; this.factory = new DefaultEventQueueFactory(); this.taskStateProvider = taskStateProvider; + this.streamLifecycleHook = streamLifecycleHook; } // For testing/extensions with custom factory and MainEventBus public InMemoryQueueManager(EventQueueFactory factory, TaskStateProvider taskStateProvider, MainEventBus mainEventBus) { + this(factory, taskStateProvider, mainEventBus, new DefaultTaskStreamLifecycleHook()); + } + + public InMemoryQueueManager(EventQueueFactory factory, TaskStateProvider taskStateProvider, + MainEventBus mainEventBus, TaskStreamLifecycleHook streamLifecycleHook) { this.factory = factory; this.taskStateProvider = taskStateProvider; this.mainEventBus = mainEventBus; + this.streamLifecycleHook = streamLifecycleHook; } @Override @@ -106,6 +120,12 @@ public EventQueue createOrTap(String taskId) { if (existing == null) { // Use builder pattern for cleaner queue creation newQueue = factory.builder(taskId).build(); + // Set the stream lifecycle hook on newly created queues only. + // If a queue already exists (putIfAbsent returns non-null), the existing + // hook is retained — the hook is bound for the lifetime of the MainQueue. + if (newQueue instanceof EventQueue.MainQueue mainQueue) { + mainQueue.setTaskStreamLifecycleHook(streamLifecycleHook); + } // Make sure an existing queue has not been added in the meantime existing = queues.putIfAbsent(taskId, newQueue); } diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/events/MainEventBusProcessor.java b/server-common/src/main/java/org/a2aproject/sdk/server/events/MainEventBusProcessor.java index eab991ff7..f98590c42 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/events/MainEventBusProcessor.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/events/MainEventBusProcessor.java @@ -257,6 +257,16 @@ private void processEvent(MainEventBusContext context) { LOGGER.debug("MainEventBusProcessor: Distributed {} to {} children for task {}", eventToDistribute.getClass().getSimpleName(), childCount, taskId); + // Step 3b: Notify stream lifecycle hook after distribution + TaskStreamLifecycleHook streamHook = mainQueue.getTaskStreamLifecycleHook(); + if (streamHook != null) { + try { + streamHook.onEvent(taskId, eventToDistribute, mainQueue.getStreamCloseHandle()); + } catch (Exception e) { + LOGGER.error("Error in TaskStreamLifecycleHook.onEvent for task {}", taskId, e); + } + } + LOGGER.debug("MainEventBusProcessor: Completed processing event for task {}", taskId); } finally { diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/events/StreamCloseHandle.java b/server-common/src/main/java/org/a2aproject/sdk/server/events/StreamCloseHandle.java new file mode 100644 index 000000000..34c8d67af --- /dev/null +++ b/server-common/src/main/java/org/a2aproject/sdk/server/events/StreamCloseHandle.java @@ -0,0 +1,25 @@ +package org.a2aproject.sdk.server.events; + +/** + * Handle for closing all active streams (ChildQueues) for a task. + * Passed to {@link TaskStreamLifecycleHook} callbacks to allow + * user-controlled stream lifecycle management. + */ +public interface StreamCloseHandle { + + /** + * Gracefully closes all active ChildQueues for this task. + * Events already in ChildQueue deques will be drained before the stream terminates. + * The MainQueue stays alive for non-finalized tasks — new clients can resubscribe, + * and the agent can keep emitting events. If the task is already finalized, the + * MainQueue will also close after the last child is removed (existing lifecycle behavior). + */ + void closeStreams(); + + /** + * Returns the number of active ChildQueues (subscribers) for this task. + * + * @return the active subscriber count + */ + int getActiveSubscriberCount(); +} diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/events/TaskStreamLifecycleHook.java b/server-common/src/main/java/org/a2aproject/sdk/server/events/TaskStreamLifecycleHook.java new file mode 100644 index 000000000..6aff27ffe --- /dev/null +++ b/server-common/src/main/java/org/a2aproject/sdk/server/events/TaskStreamLifecycleHook.java @@ -0,0 +1,77 @@ +package org.a2aproject.sdk.server.events; + +import org.a2aproject.sdk.spec.Event; + +/** + * Hook for observing task stream lifecycle events and controlling stream resources. + *

+ * Implementations are notified when clients subscribe/unsubscribe to a task's event stream + * and when events are processed for a task. The {@link StreamCloseHandle} passed to each + * callback can be used to gracefully close all active streams for the task. + *

+ * + *

Ordering guarantees

+ *
    + *
  • {@code onSubscribe} is called synchronously during {@code MainQueue.tap()}, after + * the ChildQueue has been added to the children list but before the ChildQueue is + * returned to the caller. The subscriber count visible via + * {@link StreamCloseHandle#getActiveSubscriberCount()} includes the new subscriber.
  • + *
  • {@code onEvent} is called on the {@code MainEventBusProcessor} thread after + * the event has been persisted and distributed to all ChildQueues.
  • + *
  • Because {@code onSubscribe} and {@code onEvent} run on different threads, a fast + * event emission may cause {@code onEvent} to fire concurrently with or even before + * {@code onSubscribe} returns. Stateful hook implementations must be thread-safe.
  • + *
  • {@code onUnsubscribe} is called synchronously inside {@code ChildQueue.close()}, + * on whichever thread closes the child (EventConsumer, hook via + * {@link StreamCloseHandle#closeStreams()}, or the client's transport layer).
  • + *
+ * + *

Hook binding lifetime

+ *

+ * The hook is set on a MainQueue when it is first created (in + * {@code InMemoryQueueManager.createOrTap()}). If the MainQueue already exists (e.g., a + * second client subscribes to the same task), the existing hook reference is retained. + * The hook instance is effectively bound for the lifetime of the MainQueue. + *

+ *

+ * The default implementation is a no-op. To provide custom behavior (e.g., closing streams + * after a timeout), implement this interface and register it as a CDI alternative: + *

+ *
{@code
+ * @ApplicationScoped
+ * @Alternative
+ * @Priority(1)
+ * public class MyStreamHook implements TaskStreamLifecycleHook {
+ *     // ...
+ * }
+ * }
+ * + * @see StreamCloseHandle + */ +public interface TaskStreamLifecycleHook { + + /** + * Called when a new ChildQueue is created for a task (a client subscribes to the stream). + * + * @param taskId the task identifier + * @param handle handle to close streams and query subscriber count + */ + void onSubscribe(String taskId, StreamCloseHandle handle); + + /** + * Called when a ChildQueue closes for a task (a client disconnects or streams are closed). + * + * @param taskId the task identifier + * @param handle handle to close streams and query subscriber count + */ + void onUnsubscribe(String taskId, StreamCloseHandle handle); + + /** + * Called after an event has been persisted and distributed to all ChildQueues. + * + * @param taskId the task identifier + * @param event the event that was processed + * @param handle handle to close streams and query subscriber count + */ + void onEvent(String taskId, Event event, StreamCloseHandle handle); +} diff --git a/server-common/src/test/java/org/a2aproject/sdk/server/events/EventQueueTest.java b/server-common/src/test/java/org/a2aproject/sdk/server/events/EventQueueTest.java index 3b49f6524..8504da28d 100644 --- a/server-common/src/test/java/org/a2aproject/sdk/server/events/EventQueueTest.java +++ b/server-common/src/test/java/org/a2aproject/sdk/server/events/EventQueueTest.java @@ -10,11 +10,13 @@ import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; +import java.util.ArrayList; import java.util.List; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import org.a2aproject.sdk.server.tasks.InMemoryTaskStore; +import org.a2aproject.sdk.server.tasks.MockTaskStateProvider; import org.a2aproject.sdk.server.tasks.PushNotificationSender; import org.a2aproject.sdk.spec.A2AError; import org.a2aproject.sdk.spec.Artifact; @@ -494,4 +496,102 @@ public void testMainQueueReferenceCountingStaysOpenWithActiveChildren() throws E assertTrue(mainQueue.isClosed()); assertTrue(child2.isClosed()); } + + /** + * Helper to create a queue with a TaskStateProvider that keeps MainQueue open + * when all children close (task is never finalized). + */ + private EventQueue createNonFinalizedQueueWithEventBus(String taskId) { + return EventQueueUtil.getEventQueueBuilder(mainEventBus) + .taskId(taskId) + .taskStateProvider(new MockTaskStateProvider()) + .build(); + } + + @Test + public void testCloseChildrenGracefullyClosesAllChildren() throws InterruptedException { + EventQueue mainQueue = createNonFinalizedQueueWithEventBus("close-children-test"); + + EventQueue child1 = mainQueue.tap(); + EventQueue child2 = mainQueue.tap(); + assertFalse(child1.isClosed()); + assertFalse(child2.isClosed()); + + ((EventQueue.MainQueue) mainQueue).closeChildren(); + + assertTrue(child1.isClosed()); + assertTrue(child2.isClosed()); + + assertFalse(mainQueue.isClosed()); + } + + @Test + public void testCloseChildrenAllowsNewSubscriptions() throws InterruptedException { + EventQueue mainQueue = createNonFinalizedQueueWithEventBus("close-children-resubscribe-test"); + + EventQueue child1 = mainQueue.tap(); + ((EventQueue.MainQueue) mainQueue).closeChildren(); + assertTrue(child1.isClosed()); + + EventQueue child2 = mainQueue.tap(); + assertFalse(child2.isClosed()); + assertFalse(mainQueue.isClosed()); + } + + @Test + public void testStreamCloseHandleCloseStreams() throws InterruptedException { + EventQueue mainQueue = createNonFinalizedQueueWithEventBus("handle-close-test"); + + EventQueue child1 = mainQueue.tap(); + EventQueue child2 = mainQueue.tap(); + + StreamCloseHandle handle = ((EventQueue.MainQueue) mainQueue).getStreamCloseHandle(); + assertEquals(2, handle.getActiveSubscriberCount()); + + handle.closeStreams(); + + assertTrue(child1.isClosed()); + assertTrue(child2.isClosed()); + assertFalse(mainQueue.isClosed()); + assertEquals(0, handle.getActiveSubscriberCount()); + } + + @Test + public void testStreamLifecycleHookOnEventCalledAfterDistribution() throws InterruptedException { + List receivedEvents = new ArrayList<>(); + CountDownLatch hookLatch = new CountDownLatch(1); + TaskStreamLifecycleHook hook = new TaskStreamLifecycleHook() { + @Override + public void onSubscribe(String taskId, StreamCloseHandle handle) { + } + + @Override + public void onUnsubscribe(String taskId, StreamCloseHandle handle) { + } + + @Override + public void onEvent(String taskId, Event event, StreamCloseHandle handle) { + receivedEvents.add(event); + hookLatch.countDown(); + } + }; + + EventQueue mainQueue = createQueueWithEventBus("hook-event-test"); + ((EventQueue.MainQueue) mainQueue).setTaskStreamLifecycleHook(hook); + EventQueue child = mainQueue.tap(); + + TaskStatusUpdateEvent statusEvent = TaskStatusUpdateEvent.builder() + .taskId("hook-event-test") + .contextId("ctx-1") + .status(new TaskStatus(TaskState.TASK_STATE_WORKING, null, null)) + .build(); + + mainQueue.enqueueEvent(statusEvent); + + assertTrue(hookLatch.await(5, TimeUnit.SECONDS), + "TaskStreamLifecycleHook.onEvent should have been called"); + + assertEquals(1, receivedEvents.size()); + assertSame(statusEvent, receivedEvents.get(0)); + } } diff --git a/server-common/src/test/java/org/a2aproject/sdk/server/events/InMemoryQueueManagerTest.java b/server-common/src/test/java/org/a2aproject/sdk/server/events/InMemoryQueueManagerTest.java index 71dbaa868..9fd946476 100644 --- a/server-common/src/test/java/org/a2aproject/sdk/server/events/InMemoryQueueManagerTest.java +++ b/server-common/src/test/java/org/a2aproject/sdk/server/events/InMemoryQueueManagerTest.java @@ -1,6 +1,7 @@ package org.a2aproject.sdk.server.events; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertNotSame; import static org.junit.jupiter.api.Assertions.assertNull; @@ -12,8 +13,11 @@ import java.util.List; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutionException; +import java.util.concurrent.atomic.AtomicReference; import java.util.stream.IntStream; +import org.a2aproject.sdk.spec.Event; + import org.a2aproject.sdk.server.tasks.InMemoryTaskStore; import org.a2aproject.sdk.server.tasks.MockTaskStateProvider; import org.a2aproject.sdk.server.tasks.PushNotificationSender; @@ -239,4 +243,97 @@ public void testCreateOrTapRaceCondition() throws InterruptedException, Executio long distinctCount = results.stream().distinct().count(); assertEquals(results.size(), distinctCount, "All ChildQueues should be distinct instances"); } + + @Test + public void testHookOnSubscribeCalledOnCreateOrTap() { + List subscribeEvents = new ArrayList<>(); + TaskStreamLifecycleHook hook = new TaskStreamLifecycleHook() { + @Override + public void onSubscribe(String taskId, StreamCloseHandle handle) { + subscribeEvents.add(taskId); + } + + @Override + public void onUnsubscribe(String taskId, StreamCloseHandle handle) { + } + + @Override + public void onEvent(String taskId, Event event, StreamCloseHandle handle) { + } + }; + + InMemoryQueueManager hookedQueueManager = new InMemoryQueueManager(taskStateProvider, mainEventBus, hook); + + hookedQueueManager.createOrTap("task-1"); + assertEquals(1, subscribeEvents.size()); + assertEquals("task-1", subscribeEvents.get(0)); + + hookedQueueManager.tap("task-1"); + assertEquals(2, subscribeEvents.size()); + } + + @Test + public void testHookOnUnsubscribeCalledOnChildClose() { + List unsubscribeEvents = new ArrayList<>(); + TaskStreamLifecycleHook hook = new TaskStreamLifecycleHook() { + @Override + public void onSubscribe(String taskId, StreamCloseHandle handle) { + } + + @Override + public void onUnsubscribe(String taskId, StreamCloseHandle handle) { + unsubscribeEvents.add(taskId); + } + + @Override + public void onEvent(String taskId, Event event, StreamCloseHandle handle) { + } + }; + + InMemoryQueueManager hookedQueueManager = new InMemoryQueueManager(taskStateProvider, mainEventBus, hook); + + EventQueue child = hookedQueueManager.createOrTap("task-1"); + child.close(); + + assertEquals(1, unsubscribeEvents.size()); + assertEquals("task-1", unsubscribeEvents.get(0)); + } + + @Test + public void testStreamCloseHandleClosesAllChildrenViaQueueManager() { + AtomicReference capturedHandle = new AtomicReference<>(); + TaskStreamLifecycleHook hook = new TaskStreamLifecycleHook() { + @Override + public void onSubscribe(String taskId, StreamCloseHandle handle) { + capturedHandle.set(handle); + } + + @Override + public void onUnsubscribe(String taskId, StreamCloseHandle handle) { + } + + @Override + public void onEvent(String taskId, Event event, StreamCloseHandle handle) { + } + }; + + InMemoryQueueManager hookedQueueManager = new InMemoryQueueManager(taskStateProvider, mainEventBus, hook); + + EventQueue child1 = hookedQueueManager.createOrTap("task-1"); + EventQueue child2 = hookedQueueManager.tap("task-1"); + + assertNotNull(capturedHandle.get()); + assertEquals(2, capturedHandle.get().getActiveSubscriberCount()); + + capturedHandle.get().closeStreams(); + + assertTrue(child1.isClosed()); + assertTrue(child2.isClosed()); + assertEquals(0, capturedHandle.get().getActiveSubscriberCount()); + + // MainQueue should still exist and accept new subscriptions + EventQueue child3 = hookedQueueManager.tap("task-1"); + assertNotNull(child3); + assertFalse(child3.isClosed()); + } }