Skip to content

Commit b6e2b87

Browse files
authored
fix: ensure correct task/contextId in emitted Messages (#976)
1 parent 1e61511 commit b6e2b87

2 files changed

Lines changed: 71 additions & 7 deletions

File tree

server-common/src/main/java/org/a2aproject/sdk/server/tasks/AgentEmitter.java

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66
import java.util.concurrent.atomic.AtomicBoolean;
77

88
import org.a2aproject.sdk.server.agentexecution.RequestContext;
9+
import org.slf4j.Logger;
10+
import org.slf4j.LoggerFactory;
911
import org.a2aproject.sdk.server.events.EventQueue;
1012
import org.a2aproject.sdk.spec.A2AError;
1113
import org.a2aproject.sdk.spec.Artifact;
@@ -94,6 +96,8 @@
9496
* @since 1.0.0
9597
*/
9698
public class AgentEmitter {
99+
private static final Logger LOGGER = LoggerFactory.getLogger(AgentEmitter.class);
100+
97101
private final EventQueue eventQueue;
98102
private final String taskId;
99103
private final String contextId;
@@ -508,6 +512,14 @@ public void sendMessage(List<Part<?>> parts, @Nullable Map<String, Object> metad
508512
* @since 1.0.0
509513
*/
510514
public void sendMessage(Message message) {
515+
if (message.taskId() != null && !message.taskId().equals(taskId)) {
516+
LOGGER.error("Message taskId mismatch: expected={}, actual={}", taskId, message.taskId());
517+
throw new IllegalArgumentException("Message taskId does not match the emitter's taskId");
518+
}
519+
if (message.contextId() != null && !message.contextId().equals(contextId)) {
520+
LOGGER.error("Message contextId mismatch: expected={}, actual={}", contextId, message.contextId());
521+
throw new IllegalArgumentException("Message contextId does not match the emitter's contextId");
522+
}
511523
eventQueue.enqueueEvent(message);
512524
}
513525

server-common/src/test/java/org/a2aproject/sdk/server/tasks/AgentEmitterTest.java

Lines changed: 59 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ public class AgentEmitterTest {
4444
private static final List<Part<?>> SAMPLE_PARTS = List.of(new TextPart("Test message"));
4545

4646
private static final PushNotificationSender NOOP_PUSHNOTIFICATION_SENDER = (event, snapshot) -> {};
47+
public static final int WAIT_MILLI_SECONDS = 5000;
4748

4849
EventQueue eventQueue;
4950
private MainEventBus mainEventBus;
@@ -82,7 +83,7 @@ public void cleanup() {
8283
@Test
8384
public void testAddArtifactWithCustomIdAndName() throws Exception {
8485
agentEmitter.addArtifact(SAMPLE_PARTS, "custom-artifact-id", "Custom Artifact", null);
85-
EventQueueItem item = eventQueue.dequeueEventItem(5000);
86+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
8687
assertNotNull(item);
8788
Event event = item.getEvent();
8889
assertNotNull(event);
@@ -267,7 +268,7 @@ public void testNewAgentMessageWithMetadata() throws Exception {
267268
@Test
268269
public void testAddArtifactWithAppendTrue() throws Exception {
269270
agentEmitter.addArtifact(SAMPLE_PARTS, "artifact-id", "Test Artifact", null, true, null);
270-
EventQueueItem item = eventQueue.dequeueEventItem(5000);
271+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
271272
assertNotNull(item);
272273
Event event = item.getEvent();
273274
assertNotNull(event);
@@ -288,7 +289,7 @@ public void testAddArtifactWithAppendTrue() throws Exception {
288289
@Test
289290
public void testAddArtifactWithLastChunkTrue() throws Exception {
290291
agentEmitter.addArtifact(SAMPLE_PARTS, "artifact-id", "Test Artifact", null, null, true);
291-
EventQueueItem item = eventQueue.dequeueEventItem(5000);
292+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
292293
assertNotNull(item);
293294
Event event = item.getEvent();
294295
assertNotNull(event);
@@ -305,7 +306,7 @@ public void testAddArtifactWithLastChunkTrue() throws Exception {
305306
@Test
306307
public void testAddArtifactWithAppendAndLastChunk() throws Exception {
307308
agentEmitter.addArtifact(SAMPLE_PARTS, "artifact-id", "Test Artifact", null, true, false);
308-
EventQueueItem item = eventQueue.dequeueEventItem(5000);
309+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
309310
assertNotNull(item);
310311
Event event = item.getEvent();
311312
assertNotNull(event);
@@ -321,7 +322,7 @@ public void testAddArtifactWithAppendAndLastChunk() throws Exception {
321322
@Test
322323
public void testAddArtifactGeneratesIdWhenNull() throws Exception {
323324
agentEmitter.addArtifact(SAMPLE_PARTS, null, "Test Artifact", null);
324-
EventQueueItem item = eventQueue.dequeueEventItem(5000);
325+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
325326
assertNotNull(item);
326327
Event event = item.getEvent();
327328
assertNotNull(event);
@@ -419,7 +420,7 @@ public void testConcurrentCompletionAttempts() throws Exception {
419420
thread2.join();
420421

421422
// Exactly one event should have been queued
422-
EventQueueItem item = eventQueue.dequeueEventItem(5000);
423+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
423424
assertNotNull(item);
424425
Event event = item.getEvent();
425426
assertNotNull(event);
@@ -433,9 +434,60 @@ public void testConcurrentCompletionAttempts() throws Exception {
433434
assertNull(eventQueue.dequeueEventItem(0));
434435
}
435436

437+
@Test
438+
public void sendMessageWithMatchingIdsSucceeds() throws Exception {
439+
agentEmitter.sendMessage(SAMPLE_MESSAGE);
440+
441+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
442+
assertNotNull(item);
443+
assertInstanceOf(Message.class, item.getEvent());
444+
Message message = (Message) item.getEvent();
445+
assertEquals(TEST_TASK_ID, message.taskId());
446+
assertEquals(TEST_TASK_CONTEXT_ID, message.contextId());
447+
}
448+
449+
@Test
450+
public void sendMessageWithNullIdsSucceeds() throws Exception {
451+
Message message = Message.builder()
452+
.role(ROLE_AGENT)
453+
.parts(new TextPart("no ids"))
454+
.build();
455+
agentEmitter.sendMessage(message);
456+
457+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
458+
assertNotNull(item);
459+
assertInstanceOf(Message.class, item.getEvent());
460+
}
461+
462+
@Test
463+
public void sendMessageWithMismatchedTaskIdThrows() {
464+
Message mismatchedMessage = Message.builder()
465+
.taskId("wrong-task-id")
466+
.contextId(TEST_TASK_CONTEXT_ID)
467+
.role(ROLE_AGENT)
468+
.parts(new TextPart("mismatched"))
469+
.build();
470+
IllegalArgumentException ex = assertThrows(IllegalArgumentException.class,
471+
() -> agentEmitter.sendMessage(mismatchedMessage));
472+
assertTrue(ex.getMessage().contains("Message taskId does not match"));
473+
}
474+
475+
@Test
476+
public void sendMessageWithMismatchedContextIdThrows() {
477+
Message mismatchedMessage = Message.builder()
478+
.taskId(TEST_TASK_ID)
479+
.contextId("wrong-context-id")
480+
.role(ROLE_AGENT)
481+
.parts(new TextPart("mismatched"))
482+
.build();
483+
IllegalArgumentException ex = assertThrows(IllegalArgumentException.class,
484+
() -> agentEmitter.sendMessage(mismatchedMessage));
485+
assertTrue(ex.getMessage().contains("Message contextId does not match"));
486+
}
487+
436488
private TaskStatusUpdateEvent checkTaskStatusUpdateEventOnQueue(boolean isFinal, TaskState state, Message statusMessage) throws Exception {
437489
// Wait up to 5 seconds for event (async MainEventBusProcessor needs time to distribute)
438-
EventQueueItem item = eventQueue.dequeueEventItem(5000);
490+
EventQueueItem item = eventQueue.dequeueEventItem(WAIT_MILLI_SECONDS);
439491
assertNotNull(item);
440492
Event event = item.getEvent();
441493

0 commit comments

Comments
 (0)