From abd26c796b6986ded3117c4266a9415b0b66a257 Mon Sep 17 00:00:00 2001 From: Dai Ha Date: Sat, 15 Aug 2026 05:49:27 +0200 Subject: [PATCH] CB-574: retain async task correlation --- .../dev/ltms/bridged/msg/MessageService.java | 66 +++++++++++++------ .../dev/ltms/bridged/mcp/BridgeMcpTest.java | 65 ++++++++++++++++-- 2 files changed, 106 insertions(+), 25 deletions(-) diff --git a/bridged/src/main/java/dev/ltms/bridged/msg/MessageService.java b/bridged/src/main/java/dev/ltms/bridged/msg/MessageService.java index 9feff3f..33bd7d8 100644 --- a/bridged/src/main/java/dev/ltms/bridged/msg/MessageService.java +++ b/bridged/src/main/java/dev/ltms/bridged/msg/MessageService.java @@ -155,6 +155,7 @@ public final class MessageService { private final CompletableFuture future = new CompletableFuture<>(); private final long createdNanos = System.nanoTime(); private volatile Reply question; + private volatile String turnId; private Task(String target) { this.target = target; @@ -169,8 +170,10 @@ public final class MessageService { private final Metrics metrics; // CB-502: nullable — no registry in unit tests private final ConcurrentHashMap sessionLocks = new ConcurrentHashMap<>(); private final ConcurrentHashMap tasks = new ConcurrentHashMap<>(); - /** One active async task per target, used to retain an ask after its initial send waiter resolves. */ + /** The async task that currently owns a target's send lock. */ private final ConcurrentHashMap asyncTasksByTarget = new ConcurrentHashMap<>(); + /** Async tickets paused on a specific {@code bridge_ask} turn. */ + private final ConcurrentHashMap asyncTasksByTurn = new ConcurrentHashMap<>(); private final AtomicLong ticketSeq = new AtomicLong(); private final ExecutorService asyncExecutor = Executors.newThreadPerTaskExecutor( Thread.ofVirtual().name("bridge-async-", 0).factory()); @@ -233,12 +236,6 @@ public final class MessageService { count(BridgedMetrics.REPLIES, "path", "rendezvous"); return true; // a live send took it — unchanged fast path } - Task task = asyncTasksByTarget.get(session); - if (task != null) { - finishAsyncTask(task, new Reply(Outcome.REPLIED, content)); - count(BridgedMetrics.REPLIES, "path", "rendezvous"); - return true; - } inbox.publish(session, UUID.randomUUID().toString(), content); // A rising inbox share is the signal CB-307 exists to make visible: the worker finished but // nobody was waiting, so delivery now depends on the push loop and a drain. @@ -354,6 +351,9 @@ public final class MessageService { return new Reply(Outcome.BUSY, null); // another send held the session the whole window } try { + if (hasAsyncQuestion(target)) { + return new Reply(Outcome.BUSY, null); // the worker's current turn is paused for its lead + } // Open the waiter BEFORE queueing delivery (CB-548). A fast reply — the worker already // injectable the instant we enqueue — otherwise arrives before the waiter is registered // and orphans into the inbox while this send blocks to the timeout (the enqueue-before- @@ -421,7 +421,7 @@ public final class MessageService { return new AskResult(AskOutcome.ANSWERED, answer); } catch (TimeoutException e) { log.debug("bridge_ask from {} went unanswered in {}ms", workerSession, timeoutMillis); - clearAsyncQuestion(workerSession, ticket.turnId()); + clearAsyncQuestion(ticket.turnId(), true); return new AskResult(AskOutcome.TIMED_OUT, null); } catch (ExecutionException e) { Throwable cause = e.getCause(); @@ -464,11 +464,11 @@ public final class MessageService { rendezvous.close(workerSession, reply); return new Reply(Outcome.STALE_TURN, null); // lapsed between the lookup and the unblock } - clearAsyncQuestion(workerSession, turnId); + clearAsyncQuestion(turnId, false); try { Rendezvous.Resolution r = reply.get(remainingMillis(deadlineNanos), TimeUnit.MILLISECONDS); Reply result = new Reply(outcomeOf(r.kind()), r.text(), r.turnId()); - finishAsyncTask(workerSession, result); + finishAsyncTask(turnId, result); return result; } catch (TimeoutException e) { // The worker resumed but hasn't replied yet — no completion fallback arms an answered @@ -511,11 +511,23 @@ public final class MessageService { String ticket = "task-" + ticketSeq.incrementAndGet(); Task task = new Task(target); tasks.put(ticket, task); - asyncTasksByTarget.put(target, task); asyncExecutor.submit(() -> { - Reply result = send(target, content, ASYNC_TIMEOUT_MS, onAccepted); - if (result.outcome() != Outcome.QUESTION) { - finishAsyncTask(task, result); + try { + Runnable trackingAccepted = () -> { + if (onAccepted != null) { + onAccepted.run(); + } + asyncTasksByTarget.put(target, task); + }; + Reply result = send(target, content, ASYNC_TIMEOUT_MS, trackingAccepted); + if (result.outcome() == Outcome.QUESTION) { + asyncTasksByTarget.remove(target, task); + } else { + finishAsyncTask(task, result); + } + } catch (Throwable t) { + task.future.completeExceptionally(t); + asyncTasksByTarget.remove(target, task); } }); pruneTerminalTickets(); @@ -581,14 +593,20 @@ public final class MessageService { Task task = asyncTasksByTarget.get(target); if (task != null) { task.question = new Reply(Outcome.QUESTION, text, turnId); + task.turnId = turnId; + asyncTasksByTurn.put(turnId, task); } } /** Clear an answered or lapsed question, but only when it matches the ticket's current turn. */ - private void clearAsyncQuestion(String target, String turnId) { - Task task = asyncTasksByTarget.get(target); - if (task != null && task.question != null && turnId.equals(task.question.turnId())) { + private void clearAsyncQuestion(String turnId, boolean forgetTurn) { + Task task = asyncTasksByTurn.get(turnId); + if (task != null && turnId.equals(task.turnId)) { task.question = null; + if (forgetTurn) { + asyncTasksByTurn.remove(turnId, task); + task.turnId = null; + } } } @@ -596,16 +614,24 @@ public final class MessageService { private void finishAsyncTask(Task task, Reply result) { task.future.complete(result); asyncTasksByTarget.remove(task.target, task); + if (task.turnId != null) { + asyncTasksByTurn.remove(task.turnId, task); + } } - /** Complete the active async ticket for a target, if this answer belongs to one. */ - private void finishAsyncTask(String target, Reply result) { - Task task = asyncTasksByTarget.get(target); + /** Complete the async ticket correlated to a specific answered turn. */ + private void finishAsyncTask(String turnId, Reply result) { + Task task = asyncTasksByTurn.get(turnId); if (task != null) { finishAsyncTask(task, result); } } + /** A new send must not open a waiter while an async ticket owns this worker's paused turn. */ + private boolean hasAsyncQuestion(String target) { + return asyncTasksByTurn.values().stream().anyMatch(task -> target.equals(task.target)); + } + /** Release the async executor. */ public void close() { asyncExecutor.shutdown(); diff --git a/bridged/src/test/java/dev/ltms/bridged/mcp/BridgeMcpTest.java b/bridged/src/test/java/dev/ltms/bridged/mcp/BridgeMcpTest.java index 64c8cf9..f45c8bb 100644 --- a/bridged/src/test/java/dev/ltms/bridged/mcp/BridgeMcpTest.java +++ b/bridged/src/test/java/dev/ltms/bridged/mcp/BridgeMcpTest.java @@ -193,13 +193,68 @@ class BridgeMcpTest { assertTrue(textOf(BridgeMcp.poll(messages, ticket, null)).startsWith("[pending")); BridgeMcp.reply(messages, "term_a", "finished after timeout"); - deadline = System.currentTimeMillis() + 3000; - McpSchema.CallToolResult done = BridgeMcp.poll(messages, ticket, null); - while (!"finished after timeout".equals(textOf(done)) && System.currentTimeMillis() < deadline) { + assertEquals("finished after timeout", messages.drainReplies("term_a").getFirst().content()); + } + + @Test + void asyncSendFailureDoesNotLeaveItsTicketPending() throws Exception { + String ticket = messages.sendAsync("term_a", "do it", () -> { + throw new IllegalStateException("accept failed"); + }); + + long deadline = System.currentTimeMillis() + 3000; + MessageService.TaskView view = messages.poll(ticket); + while (view.phase() == MessageService.Phase.PENDING && System.currentTimeMillis() < deadline) { Thread.sleep(5); - done = BridgeMcp.poll(messages, ticket, null); + view = messages.poll(ticket); } - assertEquals("finished after timeout", textOf(done)); + assertEquals(MessageService.Phase.FAILED, view.phase()); + assertEquals("accept failed", view.detail()); + } + + @Test + void anotherAsyncTicketCannotCaptureAReplyWhileTheFirstTicketIsAsking() throws Exception { + McpSchema.CallToolResult firstAccepted = BridgeMcp.sendAsync(messages, "term_a", "first", null, Set.of()); + String firstTicket = textOf(firstAccepted).substring(textOf(firstAccepted).indexOf("ticket=") + "ticket=".length()).trim(); + + long deadline = System.currentTimeMillis() + 3000; + while (!rendezvous.isWaiting("term_a") && System.currentTimeMillis() < deadline) { + Thread.sleep(5); + } + assertTrue(rendezvous.isWaiting("term_a")); + + CompletableFuture ask = CompletableFuture.supplyAsync( + () -> BridgeMcp.ask(messages, "term_a", "which config?", 5000L)); + MessageService.TaskView first = messages.poll(firstTicket); + deadline = System.currentTimeMillis() + 3000; + while (first.phase() != MessageService.Phase.ASKING && System.currentTimeMillis() < deadline) { + Thread.sleep(5); + first = messages.poll(firstTicket); + } + assertEquals(MessageService.Phase.ASKING, first.phase()); + String firstTurnId = first.turnId(); + + McpSchema.CallToolResult secondAccepted = BridgeMcp.sendAsync(messages, "term_a", "second", null, Set.of()); + String secondTicket = textOf(secondAccepted).substring(textOf(secondAccepted).indexOf("ticket=") + "ticket=".length()).trim(); + MessageService.TaskView second = messages.poll(secondTicket); + deadline = System.currentTimeMillis() + 3000; + while (second.phase() == MessageService.Phase.PENDING && System.currentTimeMillis() < deadline) { + Thread.sleep(5); + second = messages.poll(secondTicket); + } + assertEquals(MessageService.Phase.FAILED, second.phase()); + + BridgeMcp.reply(messages, "term_a", "late reply"); + assertEquals("late reply", messages.drainReplies("term_a").getFirst().content()); + + CompletableFuture answer = CompletableFuture.supplyAsync( + () -> BridgeMcp.answer(messages, firstTurnId, "config.yaml", 5000L)); + assertEquals("config.yaml", textOf(ask.get(6, TimeUnit.SECONDS))); + while (!rendezvous.isWaiting("term_a") && System.currentTimeMillis() < deadline) { + Thread.sleep(5); + } + BridgeMcp.reply(messages, "term_a", "done"); + assertEquals("done", textOf(answer.get(6, TimeUnit.SECONDS))); } @Test