CB-574: retain async task correlation
This commit is contained in:
@@ -155,6 +155,7 @@ public final class MessageService {
|
||||
private final CompletableFuture<Reply> 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<String, ReentrantLock> sessionLocks = new ConcurrentHashMap<>();
|
||||
private final ConcurrentHashMap<String, Task> 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<String, Task> asyncTasksByTarget = new ConcurrentHashMap<>();
|
||||
/** Async tickets paused on a specific {@code bridge_ask} turn. */
|
||||
private final ConcurrentHashMap<String, Task> 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();
|
||||
|
||||
@@ -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<McpSchema.CallToolResult> 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<McpSchema.CallToolResult> 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
|
||||
|
||||
Reference in New Issue
Block a user