diff --git a/bridged/src/main/java/dev/ltms/bridged/mcp/BridgeMcp.java b/bridged/src/main/java/dev/ltms/bridged/mcp/BridgeMcp.java index 0c960e4..a480aff 100644 --- a/bridged/src/main/java/dev/ltms/bridged/mcp/BridgeMcp.java +++ b/bridged/src/main/java/dev/ltms/bridged/mcp/BridgeMcp.java @@ -515,6 +515,9 @@ public final class BridgeMcp { ? "[done — worker finished without a structured bridge_reply; transcript tail follows]\n" + v.reply() : v.reply()); case PENDING -> text("[pending — " + v.detail() + "]"); + case ASKING -> text("[question — worker is waiting for your answer]\n" + v.reply() + + "\n\nAnswer it by calling bridge_send again with turnId=\"" + v.turnId() + + "\" and content set to your answer; the worker resumes the same turn."); case FAILED -> text("[failed — " + v.detail() + "]"); }; } 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 0d0ab44..58a0df8 100644 --- a/bridged/src/main/java/dev/ltms/bridged/msg/MessageService.java +++ b/bridged/src/main/java/dev/ltms/bridged/msg/MessageService.java @@ -127,6 +127,8 @@ public final class MessageService { public enum Phase { /** Delegated and in flight — queued for the worker or being worked. */ PENDING, + /** The worker is paused in {@code bridge_ask}; {@link TaskView#reply} and {@link TaskView#turnId} identify it. */ + ASKING, /** The worker's turn finished; {@link TaskView#reply} holds the answer. */ DONE, /** The delegation could not complete (timed out, worker gone, or busy). */ @@ -136,16 +138,28 @@ public final class MessageService { /** * A poll snapshot of an async delegation. * - * @param reply the answer when {@link #phase} is {@link Phase#DONE}, else {@code null} + * @param reply the answer when {@link #phase} is {@link Phase#DONE}, or the question when + * {@link #phase} is {@link Phase#ASKING}; otherwise {@code null} * @param replySource {@code "reply"} (structured {@code bridge_reply}) or {@code "transcript"} * (completion scrape) when {@link Phase#DONE}, else {@code null} - * @param detail a human note (live worker status while pending, or the failure reason) + * @param detail a human note (live worker status while pending, ask state, or failure reason) + * @param turnId correlation id for an {@link Phase#ASKING} ticket, else {@code null} */ - public record TaskView(String ticket, Phase phase, String reply, String replySource, String detail) { + public record TaskView(String ticket, Phase phase, String reply, String replySource, String detail, + String turnId) { } /** An in-flight or finished async delegation, keyed by its ticket. */ - private record Task(String target, CompletableFuture future, long createdNanos) { + private static final class Task { + private final String target; + 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; + } } private final AgentControl agents; @@ -156,6 +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<>(); + /** 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()); @@ -343,6 +361,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- @@ -403,12 +424,14 @@ public final class MessageService { rendezvous.closeAsk(ticket.turnId()); return new AskResult(AskOutcome.NO_WAITER, null); // no primary is blocked on this worker } + markAsyncQuestion(workerSession, question, ticket.turnId()); } try { String answer = ticket.answer().get(timeoutMillis, TimeUnit.MILLISECONDS); return new AskResult(AskOutcome.ANSWERED, answer); } catch (TimeoutException e) { log.debug("bridge_ask from {} went unanswered in {}ms", workerSession, timeoutMillis); + clearAsyncQuestion(ticket.turnId(), true); return new AskResult(AskOutcome.TIMED_OUT, null); } catch (ExecutionException e) { Throwable cause = e.getCause(); @@ -451,9 +474,12 @@ public final class MessageService { rendezvous.close(workerSession, reply); return new Reply(Outcome.STALE_TURN, null); // lapsed between the lookup and the unblock } + clearAsyncQuestion(turnId, false); try { Rendezvous.Resolution r = reply.get(remainingMillis(deadlineNanos), TimeUnit.MILLISECONDS); - return new Reply(outcomeOf(r.kind()), r.text(), r.turnId()); + Reply result = new Reply(outcomeOf(r.kind()), r.text(), r.turnId()); + finishAsyncTask(turnId, result); + return result; } catch (TimeoutException e) { // The worker resumed but hasn't replied yet — no completion fallback arms an answered // turn (it never re-entered the injector), so a silent worker rides out the window. @@ -493,9 +519,27 @@ public final class MessageService { */ public String sendAsync(String target, String content, Runnable onAccepted) { String ticket = "task-" + ticketSeq.incrementAndGet(); - CompletableFuture future = CompletableFuture.supplyAsync( - () -> send(target, content, ASYNC_TIMEOUT_MS, onAccepted), asyncExecutor); - tasks.put(ticket, new Task(target, future, System.nanoTime())); + Task task = new Task(target); + tasks.put(ticket, task); + asyncExecutor.submit(() -> { + 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(); log.debug("async send {} -> {}", ticket, target); return ticket; @@ -511,27 +555,32 @@ public final class MessageService { if (task == null) { return null; } - CompletableFuture f = task.future(); + CompletableFuture f = task.future; if (!f.isDone()) { - return new TaskView(ticket, Phase.PENDING, null, null, "worker " + liveStatus(task.target())); + Reply question = task.question; + if (question != null) { + return new TaskView(ticket, Phase.ASKING, question.text(), null, + "worker is waiting for your answer", question.turnId()); + } + return new TaskView(ticket, Phase.PENDING, null, null, "worker " + liveStatus(task.target), null); } Reply r; try { r = f.getNow(null); } catch (CompletionException | java.util.concurrent.CancellationException e) { Throwable cause = (e instanceof CompletionException ce && ce.getCause() != null) ? ce.getCause() : e; - return new TaskView(ticket, Phase.FAILED, null, null, cause.getMessage()); + return new TaskView(ticket, Phase.FAILED, null, null, cause.getMessage(), null); } if (r.completed()) { String source = r.outcome() == Outcome.REPLIED ? "reply" : "transcript"; - return new TaskView(ticket, Phase.DONE, r.text(), source, null); + return new TaskView(ticket, Phase.DONE, r.text(), source, null, null); } // A wedged worker (CB-109) carries the error context as its reason; the timeout/busy // outcomes carry none, so fall back to the outcome name. String detail = r.outcome() == Outcome.WORKER_FAILED && r.text() != null ? r.text() : "no reply — " + r.outcome().name().toLowerCase(); - return new TaskView(ticket, Phase.FAILED, null, null, detail); + return new TaskView(ticket, Phase.FAILED, null, null, detail, null); } /** Best-effort live worker status for a pending poll; never throws (a lookup error is just noise). */ @@ -546,7 +595,51 @@ public final class MessageService { /** Drop finished tickets older than the TTL so the registry cannot grow without bound. */ private void pruneTerminalTickets() { long cutoff = System.nanoTime() - TICKET_TTL_NANOS; - tasks.values().removeIf(t -> t.future().isDone() && t.createdNanos() < cutoff); + tasks.values().removeIf(t -> t.future.isDone() && t.createdNanos < cutoff); + } + + /** Record the active question for an async ticket; blocking sends have no entry and stay unchanged. */ + private void markAsyncQuestion(String target, String text, String turnId) { + 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 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; + } + } + } + + /** Complete and detach an async ticket after its worker's actual terminal reply. */ + 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 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. */ 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 b5b6b00..48d9b77 100644 --- a/bridged/src/test/java/dev/ltms/bridged/mcp/BridgeMcpTest.java +++ b/bridged/src/test/java/dev/ltms/bridged/mcp/BridgeMcpTest.java @@ -132,6 +132,131 @@ class BridgeMcpTest { assertEquals("async LGTM", textOf(polled)); } + @Test + void asyncSendSurfacesAnAskThenKeepsTheTicketForTheFinalReply() throws Exception { + McpSchema.CallToolResult accepted = BridgeMcp.sendAsync(messages, "term_a", "do it", null, Set.of()); + String ticket = textOf(accepted).substring(textOf(accepted).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)); + McpSchema.CallToolResult question = BridgeMcp.poll(messages, ticket, null); + deadline = System.currentTimeMillis() + 3000; + while (!textOf(question).contains("[question") && System.currentTimeMillis() < deadline) { + Thread.sleep(5); + question = BridgeMcp.poll(messages, ticket, null); + } + assertTrue(textOf(question).contains("which config?"), textOf(question)); + String questionText = textOf(question); + String afterTurnId = questionText.substring(questionText.indexOf("turnId=\"") + "turnId=\"".length()); + String turnId = afterTurnId.substring(0, afterTurnId.indexOf('"')); + + CompletableFuture answer = CompletableFuture.supplyAsync( + () -> BridgeMcp.answer(messages, turnId, "config.yaml", 5000L)); + assertEquals("config.yaml", textOf(ask.get(6, TimeUnit.SECONDS))); + + deadline = System.currentTimeMillis() + 3000; + while (!rendezvous.isWaiting("term_a") && System.currentTimeMillis() < deadline) { + Thread.sleep(5); + } + assertTrue(rendezvous.isWaiting("term_a")); + BridgeMcp.reply(messages, "term_a", "done"); + assertEquals("done", textOf(answer.get(6, TimeUnit.SECONDS))); + + McpSchema.CallToolResult done = BridgeMcp.poll(messages, ticket, null); + deadline = System.currentTimeMillis() + 3000; + while (!"done".equals(textOf(done)) && System.currentTimeMillis() < deadline) { + Thread.sleep(5); + done = BridgeMcp.poll(messages, ticket, null); + } + assertEquals("done", textOf(done)); + } + + @Test + void unansweredAsyncAskReturnsTheTicketToPending() throws Exception { + McpSchema.CallToolResult accepted = BridgeMcp.sendAsync(messages, "term_a", "do it", null, Set.of()); + String ticket = textOf(accepted).substring(textOf(accepted).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")); + + McpSchema.CallToolResult ask = BridgeMcp.ask(messages, "term_a", "still there?", 50L); + assertTrue(textOf(ask).contains("no answer"), textOf(ask)); + assertTrue(textOf(BridgeMcp.poll(messages, ticket, null)).startsWith("[pending")); + + BridgeMcp.reply(messages, "term_a", "finished after timeout"); + 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); + view = messages.poll(ticket); + } + 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 void pollUnknownTicketIsAnError() { McpSchema.CallToolResult res = BridgeMcp.poll(messages, "task-999", null);