From 8f6f655f631bd8089c2b6f4f87e56f04e4bd8bd9 Mon Sep 17 00:00:00 2001 From: Eric Date: Sun, 6 Sep 2026 00:22:24 +0800 Subject: [PATCH] fix: handle agent tool timeouts and cancellation as terminal outcomes --- .../chen/web/ai/AgentWebSocketHandler.java | 167 ++++++++++++---- .../web/ai/AgentWebSocketHandlerTest.java | 181 ++++++++++++++++++ 2 files changed, 314 insertions(+), 34 deletions(-) diff --git a/backend/web/src/main/java/org/jumpserver/chen/web/ai/AgentWebSocketHandler.java b/backend/web/src/main/java/org/jumpserver/chen/web/ai/AgentWebSocketHandler.java index ca7ecc4..7698a3a 100644 --- a/backend/web/src/main/java/org/jumpserver/chen/web/ai/AgentWebSocketHandler.java +++ b/backend/web/src/main/java/org/jumpserver/chen/web/ai/AgentWebSocketHandler.java @@ -10,6 +10,7 @@ import org.jumpserver.chen.framework.session.SessionManager; import org.jumpserver.chen.framework.session.impl.JMSSession; import org.jumpserver.chen.framework.ws.io.PacketIO; +import org.springframework.beans.factory.annotation.Autowired; import org.springframework.stereotype.Component; import org.springframework.web.socket.CloseStatus; import org.springframework.web.socket.TextMessage; @@ -17,7 +18,9 @@ import org.springframework.web.socket.handler.TextWebSocketHandler; import java.io.IOException; -import java.sql.SQLException; +import java.net.SocketTimeoutException; +import java.sql.SQLTimeoutException; +import java.time.Duration; import java.util.ArrayList; import java.util.LinkedHashMap; import java.util.List; @@ -25,12 +28,17 @@ import java.util.Map; import java.util.Set; import java.util.concurrent.ArrayBlockingQueue; +import java.util.concurrent.Callable; +import java.util.concurrent.CancellationException; import java.util.concurrent.ConcurrentHashMap; -import java.util.concurrent.Future; +import java.util.concurrent.ExecutionException; import java.util.concurrent.FutureTask; import java.util.concurrent.RejectedExecutionException; +import java.util.concurrent.ScheduledFuture; +import java.util.concurrent.ScheduledThreadPoolExecutor; import java.util.concurrent.ThreadPoolExecutor; import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; @Component @Slf4j @@ -46,7 +54,13 @@ public class AgentWebSocketHandler extends TextWebSocketHandler { private static final Set OPERATIONS = Set.of("generate", "explain", "repair"); private final SqlAgentToolService toolService; - private final Map> tasks = new ConcurrentHashMap<>(); + private final Map tasks = new ConcurrentHashMap<>(); + private final Duration toolTimeout; + private final ScheduledThreadPoolExecutor deadlineExecutor = new ScheduledThreadPoolExecutor(1, runnable -> { + Thread thread = new Thread(runnable, "chen-agent-deadline"); + thread.setDaemon(true); + return thread; + }); private final ThreadPoolExecutor toolExecutor = new ThreadPoolExecutor( 2, 4, @@ -61,9 +75,16 @@ public class AgentWebSocketHandler extends TextWebSocketHandler { new ThreadPoolExecutor.AbortPolicy() ); + @Autowired public AgentWebSocketHandler(SqlAgentToolService toolService) { + this(toolService, Duration.ofSeconds(60)); + } + + AgentWebSocketHandler(SqlAgentToolService toolService, Duration toolTimeout) { this.toolService = toolService; + this.toolTimeout = toolTimeout; this.toolExecutor.allowCoreThreadTimeOut(true); + this.deadlineExecutor.setRemoveOnCancelPolicy(true); } @Override @@ -132,7 +153,7 @@ private void sendManifest(WebSocketSession webSocket, JMSSession session) { new PacketIO(webSocket).sendPacket("mcp.manifest", manifest); } - private void handleToolRequest( + void handleToolRequest( WebSocketSession webSocket, String token, JMSSession session, @@ -158,8 +179,7 @@ private void handleToolRequest( } String taskKey = taskKey(webSocket, requestID); - String finalRequestID = requestID; - FutureTask future = new FutureTask<>(() -> { + ToolTask task = new ToolTask(webSocket, session, requestID, toolName, () -> { try { SessionManager.setContext(token); var resolved = toolService.resolveRequestContext(session, GSON.toJson(context), operation); @@ -167,42 +187,30 @@ private void handleToolRequest( if (resultJSON.length() > MAX_TOOL_RESULT_BYTES) { throw new IllegalStateException("Database tool result is too large"); } - sendToolResult(webSocket, session, finalRequestID, resultJSON); - } catch (IllegalArgumentException | IllegalStateException e) { - sendToolError(webSocket, session, finalRequestID, -32602, e.getMessage()); - } catch (SQLException e) { - log.warn("Chen database agent tool failed, sessionId={}, tool={}", - session.getJmsSession().getId(), toolName, e); - sendToolError(webSocket, session, finalRequestID, -32603, "Database metadata request failed"); - } catch (RuntimeException e) { - log.warn("Chen agent tool failed, sessionId={}, tool={}", - session.getJmsSession().getId(), toolName, e); - sendToolError(webSocket, session, finalRequestID, -32603, "Database tool failed"); + JsonParser.parseString(resultJSON).getAsJsonObject(); + return resultJSON; } finally { - tasks.remove(taskKey); + SessionManager.setContext(null); } - return null; }); - Future existing = tasks.putIfAbsent(taskKey, future); - if (existing != null) { + if (tasks.putIfAbsent(taskKey, task) != null) { sendToolError(webSocket, session, requestID, -32600, "Duplicate MCP request id"); return; } try { - toolExecutor.execute(future); + // Include queue time in the deadline, and retain the bounded worker pool even if JDBC ignores interruption. + task.setDeadline(deadlineExecutor.schedule(() -> task.cancelAs("timeout"), + toolTimeout.toMillis(), TimeUnit.MILLISECONDS)); + if (!task.isDone()) toolExecutor.execute(task); } catch (RejectedExecutionException e) { - tasks.remove(taskKey, future); - future.cancel(true); - throw e; + task.reject(e); } - } catch (RejectedExecutionException e) { - sendToolError(webSocket, session, requestID, -32000, "Database tool queue is full"); } catch (IllegalArgumentException e) { sendToolError(webSocket, session, requestID, -32602, e.getMessage()); } } - private void handleToolCancel(WebSocketSession webSocket, JMSSession session, JsonObject packet) { + void handleToolCancel(WebSocketSession webSocket, JMSSession session, JsonObject packet) { String requestID = ""; try { JsonObject request = rpcData(packet, session); @@ -214,8 +222,8 @@ private void handleToolCancel(WebSocketSession webSocket, JMSSession session, Js if (requestID.isBlank()) { throw new IllegalArgumentException("Invalid MCP cancellation id"); } - Future future = tasks.remove(taskKey(webSocket, requestID)); - boolean cancelled = future != null && future.cancel(true); + ToolTask task = tasks.get(taskKey(webSocket, requestID)); + boolean cancelled = task != null && task.cancelAs("cancelled"); sendRPC(webSocket, session, "mcp.cancel_result", Map.of( "jsonrpc", "2.0", "id", requestID, @@ -226,6 +234,96 @@ private void handleToolCancel(WebSocketSession webSocket, JMSSession session, Js } } + private final class ToolTask extends FutureTask { + private final WebSocketSession webSocket; + private final JMSSession session; + private final String requestID; + private final String toolName; + private volatile ScheduledFuture deadline; + private String cancellationStatus = "cancelled"; + + ToolTask(WebSocketSession webSocket, JMSSession session, String requestID, + String toolName, Callable callable) { + super(callable); + this.webSocket = webSocket; + this.session = session; + this.requestID = requestID; + this.toolName = toolName; + } + + void setDeadline(ScheduledFuture deadline) { + this.deadline = deadline; + if (isDone()) deadline.cancel(false); + } + + synchronized boolean cancelAs(String status) { + if (isDone()) return false; + cancellationStatus = status; + return cancel(true); + } + + void reject(RejectedExecutionException cause) { + setException(cause); + } + + @Override + protected void done() { + tasks.remove(taskKey(webSocket, requestID), this); + if (deadline != null) deadline.cancel(false); + toolExecutor.remove(this); + // FutureTask selects one terminal outcome; an interrupted driver returning late cannot send another result. + if (!webSocket.isOpen()) return; + try { + if (isCancelled()) { + sendOutcome(cancellationStatus); + } else { + sendToolResult(webSocket, session, requestID, get()); + } + } catch (ExecutionException e) { + Throwable cause = e.getCause(); + String outcome = expectedOutcome(cause); + if (outcome != null) { + sendOutcome(outcome); + } else if (cause instanceof RejectedExecutionException) { + sendToolError(webSocket, session, requestID, -32000, "Database tool queue is full"); + } else if (cause instanceof IllegalArgumentException || cause instanceof IllegalStateException) { + sendToolError(webSocket, session, requestID, -32602, cause.getMessage()); + } else { + log.warn("Chen database agent tool failed, sessionId={}, tool={}", + session.getJmsSession().getId(), toolName, cause); + sendToolError(webSocket, session, requestID, -32603, "Database tool failed"); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + sendOutcome("cancelled"); + } + } + + private void sendOutcome(String status) { + log.debug("Chen agent tool finished, sessionId={}, tool={}, status={}", + session.getJmsSession().getId(), toolName, status); + sendRPC(webSocket, session, "mcp.response", Map.of( + "jsonrpc", "2.0", "id", requestID, + "result", Map.of( + "isError", true, + "content", List.of(Map.of("type", "text", "text", + "timeout".equals(status) ? "Database tool timed out" : "Database tool was cancelled")), + "_meta", Map.of(AGENT_BINDING_META_KEY, Map.of("status", status, "code", "tool_" + status)) + ) + )); + } + } + + private static String expectedOutcome(Throwable cause) { + // Drivers may wrap timeout and interruption exceptions in SQLException. + for (int depth = 0; cause != null && depth < 16; depth++, cause = cause.getCause()) { + if (cause instanceof SQLTimeoutException || cause instanceof SocketTimeoutException + || cause instanceof TimeoutException) return "timeout"; + if (cause instanceof InterruptedException || cause instanceof CancellationException) return "cancelled"; + } + return null; + } + private static JsonObject rpcData(JsonObject packet, JMSSession session) { if (numberValue(packet, "version") != PROTOCOL_VERSION || !session.getJmsSession().getId().equals(stringValue(packet, "resource_session_id", 128))) { @@ -436,19 +534,20 @@ public void handleTransportError(WebSocketSession webSocket, Throwable exception private void cancelSocketTasks(WebSocketSession webSocket) { String prefix = webSocket.getId() + "\u0000"; - for (Map.Entry> entry : tasks.entrySet()) { + for (Map.Entry entry : tasks.entrySet()) { if (entry.getKey().startsWith(prefix) && tasks.remove(entry.getKey(), entry.getValue())) { - entry.getValue().cancel(true); + entry.getValue().cancelAs("cancelled"); } } } @PreDestroy public void shutdown() { - for (Future task : tasks.values()) { - task.cancel(true); + for (ToolTask task : tasks.values()) { + task.cancelAs("cancelled"); } tasks.clear(); + deadlineExecutor.shutdownNow(); toolExecutor.shutdownNow(); } diff --git a/backend/web/src/test/java/org/jumpserver/chen/web/ai/AgentWebSocketHandlerTest.java b/backend/web/src/test/java/org/jumpserver/chen/web/ai/AgentWebSocketHandlerTest.java index af304f1..8f372b6 100644 --- a/backend/web/src/test/java/org/jumpserver/chen/web/ai/AgentWebSocketHandlerTest.java +++ b/backend/web/src/test/java/org/jumpserver/chen/web/ai/AgentWebSocketHandlerTest.java @@ -1,10 +1,36 @@ package org.jumpserver.chen.web.ai; +import com.google.gson.JsonObject; +import com.google.gson.JsonParser; import org.junit.jupiter.api.Test; +import org.jumpserver.chen.framework.session.impl.JMSSession; +import org.jumpserver.wisp.Common; +import org.springframework.test.util.ReflectionTestUtils; +import org.springframework.web.socket.TextMessage; +import org.springframework.web.socket.WebSocketSession; +import java.sql.SQLException; +import java.sql.SQLTimeoutException; +import java.time.Duration; import java.util.Map; +import java.util.concurrent.BlockingQueue; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.LinkedBlockingQueue; +import java.util.concurrent.ThreadPoolExecutor; +import java.util.concurrent.TimeUnit; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.ArgumentMatchers.isNull; +import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; class AgentWebSocketHandlerTest { @Test @@ -16,4 +42,159 @@ void marksSqlProposalAsFinalResult() { assertEquals(true, ((Map) proposal.get("_meta")).get("com.jumpserver/finalResult")); } + + @Test + void deadlineSuppressesLateDriverResult() throws Exception { + assertSingleOutcomeAfterInterruption(true); + } + + @Test + void cancellationSuppressesLateDriverResult() throws Exception { + assertSingleOutcomeAfterInterruption(false); + } + + private void assertSingleOutcomeAfterInterruption(boolean timeout) throws Exception { + var fixture = new Fixture(timeout ? Duration.ofSeconds(1) : Duration.ofSeconds(60)); + var started = new CountDownLatch(1); + var release = new CountDownLatch(1); + when(fixture.service.execute(eq(fixture.session), isNull(), anyString(), anyString())).thenAnswer(call -> { + started.countDown(); + // Simulate JDBC returning successfully after ignoring cancellation. + boolean released = false; + while (!released) { + try { + released = release.await(3, TimeUnit.SECONDS); + if (!released) throw new SQLException("Test driver was not released"); + } catch (InterruptedException ignored) { + } + } + return "{}"; + }); + try { + fixture.submit(); + assertTrue(started.await(3, TimeUnit.SECONDS)); + if (!timeout) { + fixture.handler.handleToolCancel(fixture.socket, fixture.session, packet(""" + {"jsonrpc":"2.0","method":"notifications/cancelled","params":{"requestId":"request-1"}} + """)); + } + assertOutcome(fixture.responses.poll(3, TimeUnit.SECONDS), timeout ? "timeout" : "cancelled"); + if (!timeout) assertEquals("mcp.cancel_result", fixture.responses.poll(3, TimeUnit.SECONDS).get("type").getAsString()); + } finally { + release.countDown(); + fixture.close(); + } + assertTrue(fixture.responses.isEmpty(), "Late JDBC completion must not produce a second response"); + } + + @Test + void deadlineAlsoExpiresQueuedRequests() throws Exception { + var fixture = new Fixture(Duration.ofSeconds(1)); + var started = new CountDownLatch(2); + var release = new CountDownLatch(1); + var executor = (ThreadPoolExecutor) ReflectionTestUtils.getField(fixture.handler, "toolExecutor"); + assertNotNull(executor); + try { + for (int i = 0; i < 2; i++) { + executor.execute(() -> { + started.countDown(); + try { + release.await(3, TimeUnit.SECONDS); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + }); + } + assertTrue(started.await(3, TimeUnit.SECONDS)); + fixture.submit(); + assertOutcome(fixture.responses.poll(3, TimeUnit.SECONDS), "timeout"); + verifyNoInteractions(fixture.service); + assertTrue(executor.getQueue().isEmpty()); + } finally { + release.countDown(); + fixture.close(); + } + } + + @Test + void wrappedSqlTimeoutIsAnExpectedToolOutcome() throws Exception { + var fixture = new Fixture(Duration.ofSeconds(60)); + when(fixture.service.execute(eq(fixture.session), isNull(), anyString(), anyString())) + .thenThrow(new SQLException("Driver failed", new SQLTimeoutException("Query timed out"))); + try { + fixture.submit(); + assertOutcome(fixture.responses.poll(3, TimeUnit.SECONDS), "timeout"); + } finally { + fixture.close(); + } + } + + @Test + void ordinaryDatabaseFailureRemainsAnError() throws Exception { + var fixture = new Fixture(Duration.ofSeconds(60)); + when(fixture.service.execute(eq(fixture.session), isNull(), anyString(), anyString())) + .thenThrow(new SQLException("Connection failed")); + try { + fixture.submit(); + var response = fixture.responses.poll(3, TimeUnit.SECONDS); + assertNotNull(response); + assertEquals(-32603, response.getAsJsonObject("data").getAsJsonObject("error").get("code").getAsInt()); + } finally { + fixture.close(); + } + } + + private static void assertOutcome(JsonObject response, String status) { + assertNotNull(response); + var data = response.getAsJsonObject("data"); + assertFalse(data.has("error")); + var result = data.getAsJsonObject("result"); + assertTrue(result.get("isError").getAsBoolean()); + var metadata = result.getAsJsonObject("_meta").getAsJsonObject("com.jumpserver/agent"); + assertEquals(status, metadata.get("status").getAsString()); + assertEquals("tool_" + status, metadata.get("code").getAsString()); + } + + private static JsonObject packet(String data) { + var packet = JsonParser.parseString("{\"version\":1,\"resource_session_id\":\"session-1\"}").getAsJsonObject(); + packet.add("data", JsonParser.parseString(data)); + return packet; + } + + private static class Fixture { + final SqlAgentToolService service = mock(SqlAgentToolService.class); + final JMSSession session = mock(JMSSession.class); + final WebSocketSession socket = mock(WebSocketSession.class); + final BlockingQueue responses = new LinkedBlockingQueue<>(); + final AgentWebSocketHandler handler; + + Fixture(Duration timeout) throws Exception { + handler = new AgentWebSocketHandler(service, timeout); + when(session.getJmsSession()).thenReturn(Common.Session.newBuilder().setId("session-1").build()); + when(socket.getId()).thenReturn("socket-1"); + when(socket.isOpen()).thenReturn(true); + doAnswer(call -> { + responses.add(JsonParser.parseString(((TextMessage) call.getArgument(0)).getPayload()).getAsJsonObject()); + return null; + }).when(socket).sendMessage(any(TextMessage.class)); + } + + void submit() { + handler.handleToolRequest(socket, "token-1", session, packet(""" + {"jsonrpc":"2.0","id":"request-1","method":"tools/call","params":{ + "name":"inspect_schema","arguments":{},"_meta":{ + "com.jumpserver/agent":{"resource_session_id":"session-1","revision":1,"tool_call_id":"tool-1"}, + "com.jumpserver/sqlContext":{},"com.jumpserver/sqlOperation":"generate" + } + }} + """)); + } + + void close() throws InterruptedException { + handler.shutdown(); + var executor = (ThreadPoolExecutor) ReflectionTestUtils.getField(handler, "toolExecutor"); + assertNotNull(executor); + assertTrue(executor.awaitTermination(3, TimeUnit.SECONDS)); + } + } }