Use asynchronous MCP query operations
This commit is contained in:
parent
89126a2b9e
commit
aa308fd70f
|
|
@ -48,9 +48,12 @@ public class MCPAgentService {
|
|||
@Value("${app.mcp.bridge.open-timeout-seconds:30}")
|
||||
private long openTimeoutSeconds;
|
||||
|
||||
@Value("${app.mcp.bridge.query-timeout-seconds:120}")
|
||||
@Value("${app.mcp.bridge.query-timeout-seconds:600}")
|
||||
private long queryTimeoutSeconds;
|
||||
|
||||
@Value("${app.mcp.bridge.query-operation-poll-interval-ms:1000}")
|
||||
private long queryOperationPollIntervalMs;
|
||||
|
||||
@Value("${app.mcp.bridge.close-timeout-seconds:30}")
|
||||
private long closeTimeoutSeconds;
|
||||
|
||||
|
|
@ -131,7 +134,7 @@ public class MCPAgentService {
|
|||
"Unable to create MCP bridge session: " + body));
|
||||
}))
|
||||
.bodyToMono(SessionResponse.class);
|
||||
SessionResponse response = blockWithRetry(requestMono, Duration.ofSeconds(openTimeoutSeconds),
|
||||
SessionResponse response = blockWithoutRetry(requestMono, Duration.ofSeconds(openTimeoutSeconds),
|
||||
"create MCP bridge session");
|
||||
|
||||
if (response == null || response.session_id() == null || response.session_id().isBlank()) {
|
||||
|
|
@ -243,8 +246,18 @@ public class MCPAgentService {
|
|||
public String querySession(String sessionId, String prompt) {
|
||||
Objects.requireNonNull(sessionId, "sessionId cannot be null");
|
||||
Objects.requireNonNull(prompt, "prompt cannot be null");
|
||||
Mono<QueryResponse> requestMono = client().post()
|
||||
.uri("/sessions/{sessionId}/query", sessionId)
|
||||
QueryOperationResponse operation = createQueryOperation(sessionId, prompt);
|
||||
QueryOperationResponse completed = waitForQueryOperation(sessionId, operation);
|
||||
if (completed.result() == null || completed.result().result() == null) {
|
||||
throw new RuntimeException("MCP bridge query operation completed without a result");
|
||||
}
|
||||
touchSession(sessionId);
|
||||
return stringifyOperationResult(completed.result().result());
|
||||
}
|
||||
|
||||
private QueryOperationResponse createQueryOperation(String sessionId, String prompt) {
|
||||
Mono<QueryOperationResponse> requestMono = client().post()
|
||||
.uri("/sessions/{sessionId}/query-operations", sessionId)
|
||||
.contentType(MediaType.APPLICATION_JSON)
|
||||
.bodyValue(Map.of("query", prompt))
|
||||
.retrieve()
|
||||
|
|
@ -253,20 +266,101 @@ public class MCPAgentService {
|
|||
.flatMap(body -> {
|
||||
HttpStatusCode statusCode = clientResponse.statusCode();
|
||||
log.warn(
|
||||
"MCP bridge query failed: status={} sessionId={} baseUrl={} body={}",
|
||||
"MCP bridge create query operation failed: status={} sessionId={} baseUrl={} body={}",
|
||||
statusCode.value(), sessionId, mcpBridgeUrl, sanitizeForLog(body));
|
||||
return Mono.error(new RuntimeException(
|
||||
"Unable to execute MCP bridge query: " + body));
|
||||
"Unable to create MCP bridge query operation: " + body));
|
||||
}))
|
||||
.bodyToMono(QueryResponse.class);
|
||||
QueryResponse response = blockWithRetry(requestMono, Duration.ofSeconds(queryTimeoutSeconds),
|
||||
"execute MCP bridge query");
|
||||
|
||||
if (response == null || response.result() == null) {
|
||||
throw new RuntimeException("MCP bridge query returned an empty result");
|
||||
.bodyToMono(QueryOperationResponse.class);
|
||||
QueryOperationResponse operation = blockWithoutRetry(requestMono, Duration.ofSeconds(openTimeoutSeconds),
|
||||
"create MCP bridge query operation");
|
||||
if (operation == null || !StringUtils.hasText(operation.operation_id())) {
|
||||
throw new RuntimeException("MCP bridge returned an empty query operation id");
|
||||
}
|
||||
return operation;
|
||||
}
|
||||
|
||||
private QueryOperationResponse waitForQueryOperation(String sessionId, QueryOperationResponse operation) {
|
||||
long timeoutMs = Duration.ofSeconds(queryTimeoutSeconds).toMillis();
|
||||
long deadlineMs = System.currentTimeMillis() + timeoutMs;
|
||||
QueryOperationResponse current = operation;
|
||||
while (true) {
|
||||
String status = current.status() == null ? "" : current.status().trim().toLowerCase();
|
||||
switch (status) {
|
||||
case "completed":
|
||||
return current;
|
||||
case "failed", "cancelled":
|
||||
throw new RuntimeException("MCP bridge query operation " + status + ": "
|
||||
+ formatOperationError(current.error()));
|
||||
case "input-required":
|
||||
throw new RuntimeException("MCP bridge query operation requires interactive input");
|
||||
case "queued", "running":
|
||||
break;
|
||||
default:
|
||||
throw new RuntimeException("MCP bridge returned an unknown query operation status: " + current.status());
|
||||
}
|
||||
|
||||
long remainingMs = deadlineMs - System.currentTimeMillis();
|
||||
if (remainingMs <= 0) {
|
||||
throw new RuntimeException("MCP bridge query operation timed out after " + queryTimeoutSeconds
|
||||
+ " seconds (sessionId=" + sessionId + ", operationId=" + current.operation_id() + ")");
|
||||
}
|
||||
pauseBeforePolling(Math.min(Math.max(1, queryOperationPollIntervalMs), remainingMs));
|
||||
remainingMs = deadlineMs - System.currentTimeMillis();
|
||||
if (remainingMs <= 0) {
|
||||
throw new RuntimeException("MCP bridge query operation timed out after " + queryTimeoutSeconds
|
||||
+ " seconds (sessionId=" + sessionId + ", operationId=" + current.operation_id() + ")");
|
||||
}
|
||||
current = getQueryOperation(sessionId, current.operation_id(), remainingMs);
|
||||
}
|
||||
}
|
||||
|
||||
private QueryOperationResponse getQueryOperation(String sessionId, String operationId, long remainingMs) {
|
||||
Mono<QueryOperationResponse> requestMono = client().get()
|
||||
.uri("/sessions/{sessionId}/query-operations/{operationId}", sessionId, operationId)
|
||||
.retrieve()
|
||||
.onStatus(status -> status.isError(), clientResponse -> clientResponse.bodyToMono(String.class)
|
||||
.defaultIfEmpty("MCP bridge returned an error without body")
|
||||
.flatMap(body -> {
|
||||
HttpStatusCode statusCode = clientResponse.statusCode();
|
||||
log.warn(
|
||||
"MCP bridge get query operation failed: status={} sessionId={} operationId={} baseUrl={} body={}",
|
||||
statusCode.value(), sessionId, operationId, mcpBridgeUrl, sanitizeForLog(body));
|
||||
return Mono.error(new RuntimeException(
|
||||
"Unable to get MCP bridge query operation: " + body));
|
||||
}))
|
||||
.bodyToMono(QueryOperationResponse.class);
|
||||
return blockWithRetry(requestMono, Duration.ofMillis(Math.min(Duration.ofSeconds(openTimeoutSeconds).toMillis(), remainingMs)),
|
||||
"get MCP bridge query operation");
|
||||
}
|
||||
|
||||
private void pauseBeforePolling(long sleepMs) {
|
||||
try {
|
||||
Thread.sleep(sleepMs);
|
||||
} catch (InterruptedException ex) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new RuntimeException("Interrupted while waiting for MCP bridge query operation", ex);
|
||||
}
|
||||
}
|
||||
|
||||
private String formatOperationError(QueryOperationError error) {
|
||||
if (error == null) {
|
||||
return "no error details supplied";
|
||||
}
|
||||
return StringUtils.hasText(error.code())
|
||||
? error.code() + ": " + error.message()
|
||||
: String.valueOf(error.message());
|
||||
}
|
||||
|
||||
private String stringifyOperationResult(Object result) {
|
||||
if (result instanceof String text) {
|
||||
return text;
|
||||
}
|
||||
try {
|
||||
return mapper.writeValueAsString(result);
|
||||
} catch (Exception ex) {
|
||||
throw new RuntimeException("Unable to serialize MCP bridge query operation result", ex);
|
||||
}
|
||||
touchSession(sessionId);
|
||||
return response.result();
|
||||
}
|
||||
|
||||
public void closeSessionQuietly(String sessionId) {
|
||||
|
|
@ -315,10 +409,18 @@ public class MCPAgentService {
|
|||
}
|
||||
|
||||
private <T> T blockWithRetry(Mono<T> requestMono, Duration timeout, String operationLabel) {
|
||||
return block(requestMono, timeout, operationLabel, Math.max(0, retryAttempts));
|
||||
}
|
||||
|
||||
private <T> T blockWithoutRetry(Mono<T> requestMono, Duration timeout, String operationLabel) {
|
||||
return block(requestMono, timeout, operationLabel, 0);
|
||||
}
|
||||
|
||||
private <T> T block(Mono<T> requestMono, Duration timeout, String operationLabel, int attempts) {
|
||||
circuitBreaker.ensureClosed();
|
||||
try {
|
||||
T response = requestMono.timeout(timeout)
|
||||
.retry(Math.max(0, retryAttempts))
|
||||
.retry(attempts)
|
||||
.block();
|
||||
circuitBreaker.onSuccess();
|
||||
return response;
|
||||
|
|
@ -329,7 +431,7 @@ public class MCPAgentService {
|
|||
operationLabel,
|
||||
mcpBridgeUrl,
|
||||
timeout.toSeconds(),
|
||||
Math.max(0, retryAttempts),
|
||||
attempts,
|
||||
classifyFailure(rootCause),
|
||||
sanitizeForLog(rootCause.getMessage()),
|
||||
ex);
|
||||
|
|
@ -435,6 +537,13 @@ public class MCPAgentService {
|
|||
private record SessionResponse(String session_id) {
|
||||
}
|
||||
|
||||
private record QueryResponse(String result) {
|
||||
private record QueryOperationResponse(String operation_id, String session_id, String status,
|
||||
QueryOperationResult result, QueryOperationError error) {
|
||||
}
|
||||
|
||||
private record QueryOperationResult(Object result) {
|
||||
}
|
||||
|
||||
private record QueryOperationError(String code, String message) {
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -39,7 +39,8 @@ app.ollama.internal.key=${OLLAMA_INTERNAL_KEY:ollama}
|
|||
app.ollama.internal.url=${OLLAMA_INTERNAL_URL:https://ollama.internal/api}
|
||||
app.mcp.bridge.url=${MCP_BRIDGE_URL:http://localhost:8000}
|
||||
app.mcp.bridge.open-timeout-seconds=${MCP_BRIDGE_OPEN_TIMEOUT_SECONDS:90}
|
||||
app.mcp.bridge.query-timeout-seconds=${MCP_BRIDGE_QUERY_TIMEOUT_SECONDS:120}
|
||||
app.mcp.bridge.query-timeout-seconds=${MCP_BRIDGE_QUERY_TIMEOUT_SECONDS:600}
|
||||
app.mcp.bridge.query-operation-poll-interval-ms=${MCP_BRIDGE_QUERY_OPERATION_POLL_INTERVAL_MS:1000}
|
||||
app.mcp.bridge.close-timeout-seconds=${MCP_BRIDGE_CLOSE_TIMEOUT_SECONDS:30}
|
||||
app.mcp.bridge.retry-attempts=${MCP_BRIDGE_RETRY_ATTEMPTS:0}
|
||||
app.mcp.bridge.circuit-breaker.failure-threshold=${MCP_BRIDGE_CB_FAILURE_THRESHOLD:5}
|
||||
|
|
|
|||
|
|
@ -4,12 +4,19 @@ import static org.junit.jupiter.api.Assertions.assertEquals;
|
|||
import static org.junit.jupiter.api.Assertions.assertFalse;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.net.InetSocketAddress;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.test.util.ReflectionTestUtils;
|
||||
|
||||
import com.sun.net.httpserver.HttpExchange;
|
||||
import com.sun.net.httpserver.HttpServer;
|
||||
|
||||
import it.cnr.isti.workflow.manager.llms.ModelParameters;
|
||||
import org.junit.jupiter.api.io.TempDir;
|
||||
|
|
@ -36,6 +43,61 @@ class MCPAgentServiceTest {
|
|||
"""));
|
||||
}
|
||||
|
||||
@Test
|
||||
void querySessionCreatesOneAsyncOperationAndPollsUntilCompleted() throws Exception {
|
||||
AtomicInteger createCalls = new AtomicInteger();
|
||||
AtomicInteger pollCalls = new AtomicInteger();
|
||||
HttpServer bridge = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0);
|
||||
bridge.createContext("/sessions/session-1/query-operations", exchange -> {
|
||||
String path = exchange.getRequestURI().getPath();
|
||||
if ("POST".equals(exchange.getRequestMethod()) && path.endsWith("query-operations")) {
|
||||
createCalls.incrementAndGet();
|
||||
assertTrue(new String(exchange.getRequestBody().readAllBytes()).contains("write the state"));
|
||||
writeJson(exchange, 200, """
|
||||
{"operation_id":"operation-1","session_id":"session-1","status":"queued"}
|
||||
""");
|
||||
return;
|
||||
}
|
||||
if ("GET".equals(exchange.getRequestMethod()) && path.endsWith("operation-1")) {
|
||||
int call = pollCalls.incrementAndGet();
|
||||
writeJson(exchange, 200, call == 1
|
||||
? "{\"operation_id\":\"operation-1\",\"session_id\":\"session-1\",\"status\":\"running\"}"
|
||||
: """
|
||||
{"operation_id":"operation-1","session_id":"session-1","status":"completed",
|
||||
"result":{"result":"state written"}}
|
||||
""");
|
||||
return;
|
||||
}
|
||||
writeJson(exchange, 404, "{}");
|
||||
});
|
||||
bridge.start();
|
||||
try {
|
||||
MCPAgentService service = new MCPAgentService(
|
||||
WebClient.builder(),
|
||||
"http://127.0.0.1:" + bridge.getAddress().getPort(),
|
||||
new ObjectMapper(),
|
||||
providerFor("{ \"servers\": [] }"));
|
||||
ReflectionTestUtils.setField(service, "openTimeoutSeconds", 2L);
|
||||
ReflectionTestUtils.setField(service, "queryTimeoutSeconds", 2L);
|
||||
ReflectionTestUtils.setField(service, "queryOperationPollIntervalMs", 1L);
|
||||
|
||||
assertEquals("state written", service.querySession("session-1", "write the state"));
|
||||
assertEquals(1, createCalls.get());
|
||||
assertEquals(2, pollCalls.get());
|
||||
} finally {
|
||||
bridge.stop(0);
|
||||
}
|
||||
}
|
||||
|
||||
private static void writeJson(HttpExchange exchange, int status, String body) throws IOException {
|
||||
byte[] bytes = body.getBytes(java.nio.charset.StandardCharsets.UTF_8);
|
||||
exchange.getResponseHeaders().set("Content-Type", "application/json");
|
||||
exchange.sendResponseHeaders(status, bytes.length);
|
||||
try (var output = exchange.getResponseBody()) {
|
||||
output.write(bytes);
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void openSessionRequestIsUnchangedWithNoParameters() throws Exception {
|
||||
// The bridge is an external service with no contract in this repo. With nothing set the
|
||||
|
|
|
|||
Loading…
Reference in New Issue