diff --git a/bridged/src/test/java/dev/ltms/bridged/msg/MessageServiceTest.java b/bridged/src/test/java/dev/ltms/bridged/msg/MessageServiceTest.java index 5577eba..c14001b 100644 --- a/bridged/src/test/java/dev/ltms/bridged/msg/MessageServiceTest.java +++ b/bridged/src/test/java/dev/ltms/bridged/msg/MessageServiceTest.java @@ -10,6 +10,9 @@ import dev.ltms.bridged.inject.Injector; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import java.lang.reflect.Field; +import java.util.Map; +import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.TimeUnit; @@ -693,6 +696,38 @@ class MessageServiceTest { assertEquals(MessageService.Phase.DONE, awaitTicketPhase(next, MessageService.Phase.DONE).phase()); } + @Test + @SuppressWarnings("unchecked") + void asyncQuestionBelongsToTheTaskThatOwnsItsForwardWaiter() throws Exception { + String first = messages.sendAsync(T, "first task"); + awaitWaiting(); + String second = messages.sendAsync(T, "second task"); + + // Model the resolveQuestion/markAsyncQuestion race: another accepted task reached the target set. + Field tasksField = MessageService.class.getDeclaredField("tasks"); + tasksField.setAccessible(true); + Map tasks = (Map) tasksField.get(messages); + Field byTargetField = MessageService.class.getDeclaredField("asyncTasksByTarget"); + byTargetField.setAccessible(true); + Map> byTarget = (Map>) byTargetField.get(messages); + Set targetTasks = byTarget.get(T); + targetTasks.clear(); + targetTasks.add(tasks.get(second)); + + CompletableFuture ask = + CompletableFuture.supplyAsync(() -> messages.ask(T, "which config?", 5000)); + assertEquals(MessageService.Phase.ASKING, awaitTicketPhase(first, MessageService.Phase.ASKING).phase()); + assertEquals(MessageService.Phase.PENDING, messages.poll(second).phase()); + + MessageService.TaskView asking = messages.poll(first); + CompletableFuture answer = CompletableFuture.supplyAsync( + () -> messages.answer(asking.turnId(), "config.yaml", 5000)); + assertEquals("config.yaml", ask.get(5, TimeUnit.SECONDS).answer()); + awaitWaiting(); + assertTrue(rendezvous.resolve(T, "done")); + assertEquals(MessageService.Outcome.REPLIED, answer.get(5, TimeUnit.SECONDS).outcome()); + } + private void assertFailedTicket(String ticket, String reason) throws Exception { MessageService.TaskView view = awaitTicketPhase(ticket, MessageService.Phase.FAILED); assertEquals(reason, view.detail());