diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/events/EventConsumer.java b/server-common/src/main/java/org/a2aproject/sdk/server/events/EventConsumer.java index 531421aca..a4379fb09 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/events/EventConsumer.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/events/EventConsumer.java @@ -179,10 +179,14 @@ public Flow.Publisher consumeAll() { if (pollTimeoutsAfterAgentCompleted >= MAX_POLL_TIMEOUTS_AFTER_AGENT_COMPLETED) { LOGGER.debug("Agent completed with {} consecutive poll timeouts and empty queue, closing for graceful completion (queue={})", pollTimeoutsAfterAgentCompleted, System.identityHashCode(queue)); - queue.close(); - completed = true; - tube.complete(); - return; + if (queue.closeIfNotAwaitingFinalEvent()) { + completed = true; + tube.complete(); + return; + } + // A final event started while this consumer was timing out. + // Keep polling until it is distributed or its producer rolls back. + pollTimeoutsAfterAgentCompleted = 0; } else { LOGGER.debug("Agent completed but grace period active ({}/{} timeouts), continuing to poll (queue={})", pollTimeoutsAfterAgentCompleted, MAX_POLL_TIMEOUTS_AFTER_AGENT_COMPLETED, System.identityHashCode(queue)); 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 3c24a103d..c018bee5a 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 @@ -1,7 +1,9 @@ package org.a2aproject.sdk.server.events; +import java.util.HashSet; import java.util.List; import java.util.Objects; +import java.util.Set; import java.util.concurrent.BlockingQueue; import java.util.concurrent.CopyOnWriteArrayList; import java.util.concurrent.CountDownLatch; @@ -11,6 +13,7 @@ import java.util.concurrent.atomic.AtomicBoolean; import org.a2aproject.sdk.server.tasks.TaskStateProvider; +import org.a2aproject.sdk.spec.A2AError; import org.a2aproject.sdk.spec.Event; import org.a2aproject.sdk.spec.Task; import org.a2aproject.sdk.spec.TaskArtifactUpdateEvent; @@ -221,6 +224,22 @@ public void enqueueEvent(Event event) { enqueueItem(new LocalEventQueueItem(event)); } + /** + * Enqueues an event using the queue's bounded backpressure operation. + * + * @param event the event to enqueue + * @param timeout the maximum time to wait for queue capacity; must be non-negative + * @param unit the timeout unit + * @throws IllegalArgumentException if timeout is negative + * @throws UnsupportedOperationException if this queue does not support bounded enqueues + */ + public void enqueueEvent(Event event, long timeout, TimeUnit unit) { + if (timeout < 0) { + throw new IllegalArgumentException("timeout must be non-negative"); + } + enqueueItem(new LocalEventQueueItem(event), timeout, Objects.requireNonNull(unit, "unit")); + } + /** * Enqueues an event queue item for processing. *

@@ -233,6 +252,29 @@ public void enqueueEvent(Event event) { */ public abstract void enqueueItem(EventQueueItem item); + /** Timed form of {@link #enqueueItem(EventQueueItem)}. */ + public void enqueueItem(EventQueueItem item, long timeout, TimeUnit unit) { + throw new UnsupportedOperationException("This event queue does not support bounded enqueues"); + } + + /** Runtime exception used when a queue acquire was interrupted. */ + public static class EnqueueInterruptedException extends RuntimeException { + private static final long serialVersionUID = 1L; + + public EnqueueInterruptedException(InterruptedException cause) { + super("Unable to acquire the semaphore to enqueue the event", cause); + } + } + + /** Runtime exception used when a timed queue acquire expires. */ + public static class EnqueueTimeoutException extends RuntimeException { + private static final long serialVersionUID = 1L; + + public EnqueueTimeoutException() { + super("Timed out waiting for queue capacity to enqueue the event"); + } + } + /** * Enqueues an event directly to this specific queue only, bypassing the MainEventBus. *

@@ -346,6 +388,20 @@ public void clearAwaitingFinalEvent() { // Default no-op implementation - overridden by ChildQueue } + /** + * Closes this queue if it is not currently waiting for an in-flight final event. + * ChildQueue overrides this with an atomic check-and-close operation. + * + * @return true if the queue was closed, false if a final event is still expected + */ + public boolean closeIfNotAwaitingFinalEvent() { + if (isAwaitingFinalEvent()) { + return false; + } + close(); + return true; + } + /** * Closes this event queue gracefully, allowing pending events to be consumed. */ @@ -397,9 +453,37 @@ protected void doClose(boolean immediate) { return; } LOGGER.debug("Closing {} (immediate={})", this, immediate); + // onClose is invoked while the lock is still held so that subclasses can update + // their own state (e.g. immediateClose + queue.clear()) atomically with closed=true. + onClose(immediate); closed = true; } - // Subclasses handle immediate close logic (e.g., ChildQueue clears its local queue) + } + + /** + * Called from within the {@code synchronized(this)} block of {@link #doClose(boolean)}, + * before {@code closed} is set to {@code true}. Subclasses may override to perform + * state updates that must be atomic with the close. The default implementation is a no-op. + * + * @param immediate whether this is an immediate (non-draining) close + */ + protected void onClose(boolean immediate) { + // no-op — subclasses override + } + + /** + * Returns whether an event represents a final task state. + */ + static boolean isFinalEvent(Event event) { + if (event instanceof Task task) { + return task.status() != null && task.status().state() != null + && task.status().state().isFinal(); + } else if (event instanceof TaskStatusUpdateEvent statusUpdate) { + return statusUpdate.isFinal(); + } else if (event instanceof A2AError) { + return true; + } + return false; } static class MainQueue extends EventQueue { @@ -499,6 +583,18 @@ public int size() { @Override public void enqueueItem(EventQueueItem item) { + enqueueItemInternal(item, -1, null); + } + + @Override + public void enqueueItem(EventQueueItem item, long timeout, TimeUnit unit) { + if (timeout < 0) { + throw new IllegalArgumentException("timeout must be non-negative"); + } + enqueueItemInternal(item, timeout, Objects.requireNonNull(unit, "unit")); + } + + private void enqueueItemInternal(EventQueueItem item, long timeout, @Nullable TimeUnit unit) { // MainQueue must accept events even when closed to support: // 1. Late-arriving replicated events for non-finalized tasks // 2. Events enqueued during onClose callbacks (before super.doClose()) @@ -510,21 +606,42 @@ public void enqueueItem(EventQueueItem item) { // Validate event taskId matches queue taskId validateEventIds(event); - // Check if this is a final event BEFORE submitting to MainEventBus - // If it is, notify all children to expect it (so they wait for MainEventBusProcessor) - if (isFinalEvent(event)) { - LOGGER.debug("Final event detected, notifying {} children to expect it", children.size()); + boolean finalEvent = isFinalEvent(event); + + // Notify current children before waiting for capacity. A consumer may already be + // marked complete (for example, while waiting for a replicated final event) and + // otherwise close its queue before this producer is able to submit the event. + Set awaitingChildren = new HashSet<>(); + if (finalEvent) { for (ChildQueue child : children) { - child.expectFinalEvent(); + if (child.expectFinalEvent()) { + awaitingChildren.add(child); + } } } // Acquire semaphore for backpressure try { - semaphore.acquire(); + if (unit == null) { + semaphore.acquire(); + } else if (!semaphore.tryAcquire(timeout, unit)) { + awaitingChildren.forEach(ChildQueue::cancelPendingFinalEvent); + throw new EnqueueTimeoutException(); + } } catch (InterruptedException e) { + awaitingChildren.forEach(ChildQueue::cancelPendingFinalEvent); Thread.currentThread().interrupt(); - throw new RuntimeException("Unable to acquire the semaphore to enqueue the event", e); + throw new EnqueueInterruptedException(e); + } + + // Include children that subscribed while this producer was waiting for capacity. + if (finalEvent) { + for (ChildQueue child : children) { + if (!awaitingChildren.contains(child) && child.expectFinalEvent()) { + awaitingChildren.add(child); + } + } + awaitingChildren.forEach(ChildQueue::submitExpectedFinalEvent); } LOGGER.debug("Enqueued event {} {}", event instanceof Throwable ? event.toString() : event, this); @@ -541,6 +658,7 @@ public void enqueueItem(EventQueueItem item) { // Release the permit here to avoid leaking it and eventually blocking // all event processing for this task. semaphore.release(); + awaitingChildren.forEach(ChildQueue::cancelSubmittedFinalEvent); throw e; } } @@ -586,19 +704,6 @@ private void validateEventIds(Event event) { } } - /** - * Checks if an event represents a final task state. - */ - private boolean isFinalEvent(Event event) { - if (event instanceof Task task) { - return task.status() != null && task.status().state() != null - && task.status().state().isFinal(); - } else if (event instanceof TaskStatusUpdateEvent statusUpdate) { - return statusUpdate.isFinal(); - } - return false; - } - @Override public void awaitQueuePollerStart() throws InterruptedException { LOGGER.debug("Waiting for queue poller to start on {}", this); @@ -664,6 +769,18 @@ void distributeToChildren(EventQueueItem item) { taskId, item.getEvent().getClass().getSimpleName(), childCount); } children.forEach(child -> { + // Skip children that have already been closed. CopyOnWriteArrayList gives + // snapshot semantics, so a child removed by closeIfNotAwaitingFinalEvent() + // may still appear in this iteration. Checking isClosed() here narrows the + // TOCTOU window to the time between the check and internalEnqueueItem(); + // a narrow residual race remains for graceful close but is harmless for + // non-final events (they'd land in the draining queue) and impossible for + // final events (submittedFinalEvents prevents close while one is in flight). + if (child.isClosed()) { + LOGGER.debug("MainQueue[{}]: Skipping closed child queue for event {}", + taskId, item.getEvent().getClass().getSimpleName()); + return; + } LOGGER.debug("MainQueue[{}]: Enqueueing event {} to child queue", taskId, item.getEvent().getClass().getSimpleName()); child.internalEnqueueItem(item); @@ -776,7 +893,8 @@ static class ChildQueue extends EventQueue { private final MainQueue parent; private final BlockingQueue queue; private volatile boolean immediateClose = false; - private volatile boolean awaitingFinalEvent = false; + private int pendingFinalEvents; + private int submittedFinalEvents; public ChildQueue(MainQueue parent) { this.parent = parent; @@ -797,6 +915,11 @@ public void enqueueItem(EventQueueItem item) { parent.enqueueItem(item); } + @Override + public void enqueueItem(EventQueueItem item, long timeout, TimeUnit unit) { + parent.enqueueItem(item, timeout, unit); + } + private void internalEnqueueItem(EventQueueItem item) { // Internal method called by MainEventBusProcessor to add to local queue // Note: Semaphore is managed by parent MainQueue (acquire/release), not ChildQueue @@ -813,27 +936,18 @@ private void internalEnqueueItem(EventQueueItem item) { } else { LOGGER.debug("Enqueued event {} {}", event instanceof Throwable ? event.toString() : event, this); - // If we were awaiting a final event and this is it, clear the flag - if (awaitingFinalEvent && isFinalEvent(event)) { - awaitingFinalEvent = false; - LOGGER.debug("ChildQueue {} received awaited final event", System.identityHashCode(this)); + // If we were awaiting a final event and this is it, decrement the counter + if (isFinalEvent(event)) { + synchronized (this) { + if (submittedFinalEvents > 0) { + submittedFinalEvents--; + LOGGER.debug("ChildQueue {} received awaited final event", System.identityHashCode(this)); + } + } } } } - /** - * Checks if an event represents a final task state. - */ - private boolean isFinalEvent(Event event) { - if (event instanceof Task task) { - return task.status() != null && task.status().state() != null - && task.status().state().isFinal(); - } else if (event instanceof TaskStatusUpdateEvent statusUpdate) { - return statusUpdate.isFinal(); - } - return false; - } - @Override public void enqueueLocalOnly(EventQueueItem item) { internalEnqueueItem(item); @@ -849,7 +963,7 @@ public EventQueueItem dequeueEventItem(int waitMilliSeconds) throws EventQueueCl // For immediate close: exit immediately even if queue is not empty (race with MainEventBusProcessor) // For graceful close: only exit when queue is empty (wait for all events to be consumed) // BUT: if awaiting final event, keep polling even if closed and empty - if (isClosed() && (queue.isEmpty() || immediateClose) && !awaitingFinalEvent) { + if (isClosed() && (queue.isEmpty() || immediateClose) && !isAwaitingFinalEvent()) { LOGGER.debug("ChildQueue is closed{}, sending termination message. {} (queueSize={})", immediateClose ? " (immediate)" : " and empty", this, @@ -893,8 +1007,8 @@ public int size() { } @Override - public boolean isAwaitingFinalEvent() { - return awaitingFinalEvent; + public synchronized boolean isAwaitingFinalEvent() { + return pendingFinalEvents > 0 || submittedFinalEvents > 0; } @Override @@ -908,16 +1022,18 @@ public void signalQueuePollerStarted() { } @Override - protected void doClose(boolean immediate) { - super.doClose(immediate); // Sets closed flag + protected void onClose(boolean immediate) { + // Invoked by EventQueue.doClose() while synchronized(this) is held, so closed, + // immediateClose, and the queue clear all become visible atomically. if (immediate) { - // Immediate close: clear pending events from local queue this.immediateClose = true; + pendingFinalEvents = 0; + submittedFinalEvents = 0; int clearedCount = queue.size(); queue.clear(); LOGGER.debug("Cleared {} events from ChildQueue for immediate close: {}", clearedCount, this); } - // For graceful close, let the queue drain naturally through normal consumption + // For graceful close the queue drains naturally; nothing to do here. } /** @@ -925,9 +1041,34 @@ protected void doClose(boolean immediate) { * Called by MainQueue when it enqueues a final event, BEFORE submitting to MainEventBus. * This ensures the ChildQueue keeps polling until the final event arrives (after MainEventBusProcessor). */ - void expectFinalEvent() { - awaitingFinalEvent = true; + synchronized boolean expectFinalEvent() { + if (isClosed()) { + return false; + } + pendingFinalEvents++; LOGGER.debug("ChildQueue {} now awaiting final event", System.identityHashCode(this)); + return true; + } + + synchronized void submitExpectedFinalEvent() { + if (pendingFinalEvents > 0) { + pendingFinalEvents--; + if (!isClosed()) { + submittedFinalEvents++; + } + } + } + + synchronized void cancelPendingFinalEvent() { + if (pendingFinalEvents > 0) { + pendingFinalEvents--; + } + } + + synchronized void cancelSubmittedFinalEvent() { + if (submittedFinalEvents > 0) { + submittedFinalEvents--; + } } /** @@ -935,11 +1076,24 @@ void expectFinalEvent() { * This allows normal timeout logic to proceed if the final event never arrives. */ @Override - public void clearAwaitingFinalEvent() { - awaitingFinalEvent = false; + public synchronized void clearAwaitingFinalEvent() { + pendingFinalEvents = 0; + submittedFinalEvents = 0; LOGGER.debug("ChildQueue {} cleared awaitingFinalEvent flag (timeout)", System.identityHashCode(this)); } + @Override + public boolean closeIfNotAwaitingFinalEvent() { + synchronized (this) { + if (pendingFinalEvents > 0 || submittedFinalEvents > 0) { + return false; + } + doClose(false); + } + parent.childClosing(this, false); + return true; + } + @Override public void close() { close(false); diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandler.java b/server-common/src/main/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandler.java index 4388b0443..f9958ad62 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandler.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandler.java @@ -187,6 +187,9 @@ */ @ApplicationScoped public class DefaultRequestHandler implements RequestHandler { + private static final int MAX_INTERRUPTED_ERROR_ENQUEUE_RETRIES = 1; + private static final long ERROR_ENQUEUE_TIMEOUT_SECONDS = 5; + private static final Logger LOGGER = LoggerFactory.getLogger(DefaultRequestHandler.class); @@ -1255,22 +1258,29 @@ public void run() { LOGGER.debug("Agent execution starting for task {}", taskId); AgentEmitter emitter = new AgentEmitter(requestContext, queue); try { - resolvedExecutor.execute(requestContext, emitter); - } catch (A2AError e) { - // Log A2A errors at WARN level with full stack trace - // These are expected business errors but should be tracked - LOGGER.warn("Agent execution threw A2AError for task {}: {} - {}", - taskId, e.getClass().getSimpleName(), e.getMessage(), e); - emitter.fail(e); - } catch (RuntimeException e) { - // Log unexpected runtime exceptions at ERROR level - // These indicate bugs in agent implementation - LOGGER.error("Agent execution threw unexpected RuntimeException for task {}", taskId, e); - emitter.fail(new org.a2aproject.sdk.spec.InternalError("Agent execution failed: " + e.getMessage())); - } catch (Exception e) { - // Log other exceptions at ERROR level - LOGGER.error("Agent execution threw unexpected Exception for task {}", taskId, e); - emitter.fail(new org.a2aproject.sdk.spec.InternalError("Agent execution failed: " + e.getMessage())); + try { + resolvedExecutor.execute(requestContext, emitter); + } catch (A2AError e) { + // Log A2A errors at WARN level with full stack trace + // These are expected business errors but should be tracked + LOGGER.warn("Agent execution threw A2AError for task {}: {} - {}", + taskId, e.getClass().getSimpleName(), e.getMessage(), e); + enqueueErrorPreservingInterrupt(emitter, e); + } catch (RuntimeException e) { + // Log unexpected runtime exceptions at ERROR level + // These indicate bugs in agent implementation + LOGGER.error("Agent execution threw unexpected RuntimeException for task {}", taskId, e); + enqueueErrorPreservingInterrupt(emitter, + new InternalError("Agent execution failed: " + e.getMessage())); + } catch (Exception e) { + // Log other exceptions at ERROR level + LOGGER.error("Agent execution threw unexpected Exception for task {}", taskId, e); + enqueueErrorPreservingInterrupt(emitter, + new InternalError("Agent execution failed: " + e.getMessage())); + } + } finally { + // Executor threads are reused. Do not return an agent's interrupt to the pool. + Thread.interrupted(); } LOGGER.debug("Agent execution completed for task {}", taskId); // The consumer (running on the Vert.x worker thread) handles queue lifecycle. @@ -1310,6 +1320,40 @@ public void run() { return runnable; } + private void enqueueErrorPreservingInterrupt(AgentEmitter emitter, A2AError error) { + boolean wasInterrupted = Thread.interrupted(); + int retries = 0; + try { + while (true) { + try { + emitter.tryFailWithTimeout(error, ERROR_ENQUEUE_TIMEOUT_SECONDS, SECONDS); + return; + } catch (EventQueue.EnqueueTimeoutException e) { + // Queue stayed full for the entire timeout window — the infrastructure is + // under severe backpressure and we cannot deliver the terminal error event. + // Log and return so the agent thread exits cleanly; the EventConsumer will + // eventually close the stream via its normal timeout path. + LOGGER.warn("Timed out after {} s enqueueing error event for task after {} attempt(s) — task may not receive terminal status", + ERROR_ENQUEUE_TIMEOUT_SECONDS, retries + 1); + return; + } catch (EventQueue.EnqueueInterruptedException e) { + wasInterrupted = true; + if (retries >= MAX_INTERRUPTED_ERROR_ENQUEUE_RETRIES) { + LOGGER.warn("Interrupted {} time(s) enqueueing error event for task — task may not receive terminal status", + retries + 1); + return; + } + retries++; + wasInterrupted |= Thread.interrupted(); + } + } + } finally { + if (wasInterrupted) { + Thread.currentThread().interrupt(); + } + } + } + private CompletableFuture cleanupProducer(@Nullable CompletableFuture agentFuture, @Nullable CompletableFuture consumptionFuture, String taskId, EventQueue queue, boolean isStreaming) { LOGGER.debug("Starting cleanup for task {} (streaming={})", taskId, isStreaming); logThreadStats("CLEANUP START"); diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/tasks/AgentEmitter.java b/server-common/src/main/java/org/a2aproject/sdk/server/tasks/AgentEmitter.java index 19d458fd6..2f5dcad8a 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/tasks/AgentEmitter.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/tasks/AgentEmitter.java @@ -3,6 +3,7 @@ import java.util.List; import java.util.Map; import java.util.UUID; +import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import org.a2aproject.sdk.server.agentexecution.RequestContext; @@ -128,11 +129,9 @@ public void updateStatus(TaskState taskState, @Nullable Message message) { throw new IllegalStateException("Cannot update task status - terminal state already reached"); } - // For final states, atomically set the flag - if (isFinal) { - if (!terminalStateReached.compareAndSet(false, true)) { - throw new IllegalStateException("Cannot update task status - terminal state already reached"); - } + // Claim the terminal transition atomically without holding a lock during queue backpressure. + if (isFinal && !terminalStateReached.compareAndSet(false, true)) { + throw new IllegalStateException("Cannot update task status - terminal state already reached"); } TaskStatusUpdateEvent event = TaskStatusUpdateEvent.builder() @@ -140,7 +139,29 @@ public void updateStatus(TaskState taskState, @Nullable Message message) { .contextId(contextId) .status(new TaskStatus(taskState, message, null)) .build(); - eventQueue.enqueueEvent(event); + if (isFinal) { + enqueueTerminalEvent(event); + } else { + eventQueue.enqueueEvent(event); + } + } + + private void enqueueTerminalEvent(Event event) { + try { + eventQueue.enqueueEvent(event); + } catch (RuntimeException | Error e) { + terminalStateReached.compareAndSet(true, false); + throw e; + } + } + + private void enqueueTerminalEvent(Event event, long timeout, TimeUnit unit) { + try { + eventQueue.enqueueEvent(event, timeout, unit); + } catch (RuntimeException | Error e) { + terminalStateReached.compareAndSet(true, false); + throw e; + } } /** @@ -278,17 +299,44 @@ public void fail(@Nullable Message message) { * @since 1.0.0 */ public void fail(A2AError error) { - // Set terminal state flag BEFORE enqueueing error - // This prevents race conditions where agent calls fail(error) then complete() - if (!terminalStateReached.compareAndSet(false, true)) { + if (!tryFail(error)) { throw new IllegalStateException("Cannot update task status - terminal state already reached"); } - - eventQueue.enqueueEvent(error); // Status transition happens automatically in MainEventBusProcessor // The error event is terminal and will trigger FAILED state transition } + /** + * Attempts to fail the task unless a terminal state has already been claimed. + * + * @param error the A2A error to enqueue + * @return {@code true} if the error was enqueued, or {@code false} if a terminal state was already reached + */ + public boolean tryFail(A2AError error) { + if (!terminalStateReached.compareAndSet(false, true)) { + return false; + } + enqueueTerminalEvent(error); + // Status transition happens automatically in MainEventBusProcessor + return true; + } + + /** + * Attempts to fail the task, waiting no longer than the supplied duration for queue capacity. + * + * @param error the A2A error to enqueue + * @param timeout maximum time to wait for queue capacity + * @param unit timeout unit + * @return {@code true} if the error was enqueued, or {@code false} if a terminal state was already reached + */ + public boolean tryFailWithTimeout(A2AError error, long timeout, TimeUnit unit) { + if (!terminalStateReached.compareAndSet(false, true)) { + return false; + } + enqueueTerminalEvent(error, timeout, unit); + return true; + } + /** * Marks the task as SUBMITTED. */ diff --git a/server-common/src/test/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandlerTest.java b/server-common/src/test/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandlerTest.java index 1e729bd20..95ffd0f7c 100644 --- a/server-common/src/test/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandlerTest.java +++ b/server-common/src/test/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandlerTest.java @@ -10,6 +10,7 @@ import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; +import java.util.Arrays; import java.util.List; import java.util.Map; import java.util.Set; @@ -35,17 +36,20 @@ import org.a2aproject.sdk.server.events.InMemoryQueueManager; import org.a2aproject.sdk.server.events.MainEventBus; import org.a2aproject.sdk.server.events.MainEventBusProcessor; +import org.a2aproject.sdk.server.events.MainEventBusProcessorCallback; import org.a2aproject.sdk.server.tasks.AgentEmitter; import org.a2aproject.sdk.server.tasks.InMemoryPushNotificationConfigStore; import org.a2aproject.sdk.server.tasks.InMemoryTaskStore; import org.a2aproject.sdk.server.tasks.PushNotificationConfigStore; import org.a2aproject.sdk.server.tasks.PushNotificationSender; import org.a2aproject.sdk.server.tasks.TaskStore; +import org.a2aproject.sdk.server.tasks.TaskStateProvider; import org.a2aproject.sdk.spec.A2AError; import org.a2aproject.sdk.spec.CancelTaskParams; import org.a2aproject.sdk.spec.Event; import org.a2aproject.sdk.spec.EventKind; import org.a2aproject.sdk.spec.GetTaskPushNotificationConfigParams; +import org.a2aproject.sdk.spec.InternalError; import org.a2aproject.sdk.spec.InvalidParamsError; import org.a2aproject.sdk.spec.ListTasksParams; import org.a2aproject.sdk.spec.ListTaskPushNotificationConfigsParams; @@ -1692,6 +1696,185 @@ public void cancel(RequestContext context, AgentEmitter emitter) { assertEquals(TaskState.TASK_STATE_FAILED, storedTask.status().state()); } + @Test + void interruptedTerminalEnqueueStillReportsAgentFailure() throws Exception { + assertInterruptedTerminalEnqueueReportsError( + (context, emitter) -> emitter.complete(), InternalError.class, null, true, false); + } + + @Test + void interruptedProtocolErrorEnqueueStillReportsAgentFailure() throws Exception { + CountDownLatch agentReadyToBeInterrupted = new CountDownLatch(1); + CountDownLatch agentWait = new CountDownLatch(1); + try { + assertInterruptedTerminalEnqueueReportsError((context, emitter) -> { + agentReadyToBeInterrupted.countDown(); + try { + agentWait.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new UnsupportedOperationError(); + } + }, UnsupportedOperationError.class, agentReadyToBeInterrupted, false, true); + } finally { + agentWait.countDown(); + } + } + + private void assertInterruptedTerminalEnqueueReportsError( + AgentExecutorMethod agentAction, + Class expectedErrorClass, + CountDownLatch agentReadyToBeInterrupted, + boolean interruptBlockedTerminalEnqueue, + boolean returnImmediately) throws Exception { + EventQueueUtil.stop(mainEventBusProcessor); + queueManager = new InMemoryQueueManager((TaskStateProvider) taskStore, mainEventBus) { + @Override + public EventQueue.EventQueueBuilder createBaseEventQueueBuilder(String taskId) { + return super.createBaseEventQueueBuilder(taskId).queueSize(1); + } + }; + mainEventBusProcessor = new MainEventBusProcessor(mainEventBus, taskStore, + NOOP_PUSHNOTIFICATION_SENDER, queueManager); + EventQueueUtil.start(mainEventBusProcessor); + + CountDownLatch workingEventProcessing = new CountDownLatch(1); + CountDownLatch releaseWorkingEvent = new CountDownLatch(1); + CountDownLatch reportedErrorProcessed = new CountDownLatch(1); + AtomicReference reportedError = new AtomicReference<>(); + mainEventBusProcessor.setCallback(new MainEventBusProcessorCallback() { + @Override + public void onEventProcessed(String taskId, Event event) { + if (event instanceof TaskStatusUpdateEvent statusUpdate + && statusUpdate.status().state() == TaskState.TASK_STATE_WORKING) { + workingEventProcessing.countDown(); + try { + releaseWorkingEvent.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } else if (event instanceof A2AError error && expectedErrorClass.isInstance(error)) { + reportedError.set(error); + reportedErrorProcessed.countDown(); + } + } + + @Override + public void onTaskFinalized(String taskId) { + } + }); + + CountDownLatch agentFinished = new CountDownLatch(1); + AtomicReference agentThread = new AtomicReference<>(); + AtomicReference interruptStatusAtPoolReturn = new AtomicReference<>(false); + Executor recordingExecutor = command -> internalExecutor.execute(() -> { + try { + command.run(); + } finally { + interruptStatusAtPoolReturn.set(Thread.currentThread().isInterrupted()); + Thread.interrupted(); + agentFinished.countDown(); + } + }); + requestHandler = DefaultRequestHandler.builder() + .agentExecutor(executor) + .taskStore(taskStore) + .queueManager(queueManager) + .pushConfigStore(pushConfigStore) + .pushNotificationsEnabled(true) + .mainEventBusProcessor(mainEventBusProcessor) + .executor(recordingExecutor) + .eventConsumerExecutor(internalExecutor) + .authorizationRequired(false) + .build(); + + CountDownLatch executorRan = new CountDownLatch(1); + agentExecutorExecute = (context, emitter) -> { + agentThread.set(Thread.currentThread()); + executorRan.countDown(); + emitter.startWork(); + agentAction.invoke(context, emitter); + }; + + MessageSendParams params = MessageSendParams.builder() + .message(Message.builder() + .messageId("msg-interrupted-terminal-enqueue") + .role(Message.Role.ROLE_USER) + .parts(new TextPart("hello")) + .build()) + .configuration(MessageSendConfiguration.builder() + .returnImmediately(returnImmediately) + .acceptedOutputModes(List.of()) + .build()) + .build(); + + ExecutorService controller = Executors.newSingleThreadExecutor(); + try { + Future poolInterruptStatus = controller.submit(() -> { + assertTrue(workingEventProcessing.await(5, TimeUnit.SECONDS), + "Processor should hold the queue capacity after the WORKING event"); + assertTrue(executorRan.await(5, TimeUnit.SECONDS), "Executor should have run"); + Thread worker = agentThread.get(); + assertNotNull(worker); + if (interruptBlockedTerminalEnqueue) { + assertTrue(awaitThreadWaiting(worker, false), + "Terminal enqueue should wait while the queue is full"); + } else { + assertNotNull(agentReadyToBeInterrupted); + assertTrue(agentReadyToBeInterrupted.await(5, TimeUnit.SECONDS), + "Agent should be ready to throw the protocol error"); + } + + worker.interrupt(); + + assertTrue(awaitThreadWaiting(worker, false), + "Error enqueue should wait with the interrupt temporarily cleared"); + releaseWorkingEvent.countDown(); + assertTrue(agentFinished.await(5, TimeUnit.SECONDS), "Agent error reporting should finish"); + return interruptStatusAtPoolReturn.get(); + }); + + String taskId = ""; + if (returnImmediately) { + Task task = assertInstanceOf(Task.class, requestHandler.onMessageSend(params, NULL_CONTEXT)); + taskId = task.id(); + } else { + assertThrows(expectedErrorClass, + () -> requestHandler.onMessageSend(params, NULL_CONTEXT)); + } + + assertTrue(reportedErrorProcessed.await(5, TimeUnit.SECONDS), + "The agent error should be processed by the event bus"); + assertInstanceOf(expectedErrorClass, reportedError.get()); + if (returnImmediately) { + Task storedTask = taskStore.get(taskId); + assertNotNull(storedTask); + assertEquals(TaskState.TASK_STATE_FAILED, storedTask.status().state()); + } + assertFalse(poolInterruptStatus.get(5, TimeUnit.SECONDS), + "The agent worker should return to the executor pool without an interrupt flag"); + } finally { + releaseWorkingEvent.countDown(); + controller.shutdownNow(); + } + } + + private static boolean awaitThreadWaiting(Thread thread, boolean interrupted) throws InterruptedException { + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(30); + do { + boolean waitingOrAcquiring = thread.getState() == Thread.State.WAITING + || Arrays.stream(thread.getStackTrace()).anyMatch(frame -> + frame.getClassName().equals("java.util.concurrent.Semaphore") + && (frame.getMethodName().equals("acquire") + || frame.getMethodName().equals("tryAcquire"))); + if (waitingOrAcquiring && thread.isInterrupted() == interrupted) { + return true; + } + TimeUnit.MILLISECONDS.sleep(10); + } while (System.nanoTime() < deadline); + return false; + } + private DefaultRequestHandler buildHandlerWithRouter(AgentExecutorRouter router) { return DefaultRequestHandler.builder() .agentExecutor(executor) diff --git a/server-common/src/test/java/org/a2aproject/sdk/server/tasks/AgentEmitterConcurrencyTest.java b/server-common/src/test/java/org/a2aproject/sdk/server/tasks/AgentEmitterConcurrencyTest.java index 92745c465..832501b7e 100644 --- a/server-common/src/test/java/org/a2aproject/sdk/server/tasks/AgentEmitterConcurrencyTest.java +++ b/server-common/src/test/java/org/a2aproject/sdk/server/tasks/AgentEmitterConcurrencyTest.java @@ -2,6 +2,7 @@ import org.a2aproject.sdk.server.agentexecution.RequestContext; import org.a2aproject.sdk.server.events.EventQueue; +import org.a2aproject.sdk.spec.InternalError; import org.a2aproject.sdk.spec.UnsupportedOperationError; import org.junit.jupiter.api.Test; @@ -219,6 +220,43 @@ public void testFailWithErrorThenFailWithMessage() { exception.getMessage()); } + @Test + public void testFailedTerminalStatusEnqueueAllowsErrorFallback() { + RequestContext context = mock(RequestContext.class); + when(context.getTaskId()).thenReturn("test-task-123"); + when(context.getContextId()).thenReturn("test-context-456"); + + EventQueue eventQueue = mock(EventQueue.class); + RuntimeException enqueueFailure = new RuntimeException("queue failure"); + doThrow(enqueueFailure).doNothing().when(eventQueue).enqueueEvent(any()); + AgentEmitter emitter = new AgentEmitter(context, eventQueue); + + assertSame(enqueueFailure, assertThrows(RuntimeException.class, emitter::complete)); + + emitter.fail(new InternalError("fallback")); + + verify(eventQueue, times(2)).enqueueEvent(any()); + } + + @Test + public void testFailedErrorEnqueueAllowsTerminalStatusRetry() { + RequestContext context = mock(RequestContext.class); + when(context.getTaskId()).thenReturn("test-task-123"); + when(context.getContextId()).thenReturn("test-context-456"); + + EventQueue eventQueue = mock(EventQueue.class); + RuntimeException enqueueFailure = new RuntimeException("queue failure"); + doThrow(enqueueFailure).doNothing().when(eventQueue).enqueueEvent(any()); + AgentEmitter emitter = new AgentEmitter(context, eventQueue); + + assertSame(enqueueFailure, assertThrows(RuntimeException.class, + () -> emitter.fail(new UnsupportedOperationError()))); + + emitter.complete(); + + verify(eventQueue, times(2)).enqueueEvent(any()); + } + @Test public void testNonTerminalThenTerminalState() throws InterruptedException { // Setup