From 58ed5ea59ff4488cddb26a775c36ca2016a95526 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Mon, 10 Mar 2025 15:04:22 +0100 Subject: [PATCH] feat(mcp): handle ToolContext in MCP tool callbacks - Implement the new ToolCallback.call(String, ToolContext) method in both Sync and Async MCP tool callbacks. - Since MCP tools don't support tool context, the implementation ignores the context parameter and delegates to the existing call(String) method. Added test to verify the behavior. Resolves #2378 Signed-off-by: Christian Tzolov --- .../ai/mcp/AsyncMcpToolCallback.java | 7 +++++++ .../ai/mcp/SyncMcpToolCallback.java | 7 +++++++ .../ai/mcp/SyncMcpToolCallbackTests.java | 20 +++++++++++++++++++ 3 files changed, 34 insertions(+) diff --git a/mcp/common/src/main/java/org/springframework/ai/mcp/AsyncMcpToolCallback.java b/mcp/common/src/main/java/org/springframework/ai/mcp/AsyncMcpToolCallback.java index 121e1d9b8..a4900071d 100644 --- a/mcp/common/src/main/java/org/springframework/ai/mcp/AsyncMcpToolCallback.java +++ b/mcp/common/src/main/java/org/springframework/ai/mcp/AsyncMcpToolCallback.java @@ -22,6 +22,7 @@ import io.modelcontextprotocol.client.McpAsyncClient; import io.modelcontextprotocol.spec.McpSchema.CallToolRequest; import io.modelcontextprotocol.spec.McpSchema.Tool; +import org.springframework.ai.chat.model.ToolContext; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.tool.ToolCallback; import org.springframework.ai.tool.definition.ToolDefinition; @@ -110,4 +111,10 @@ public class AsyncMcpToolCallback implements ToolCallback { .block(); } + @Override + public String call(String toolArguments, ToolContext toolContext) { + // ToolContext is not supported by the MCP tools + return this.call(toolArguments); + } + } diff --git a/mcp/common/src/main/java/org/springframework/ai/mcp/SyncMcpToolCallback.java b/mcp/common/src/main/java/org/springframework/ai/mcp/SyncMcpToolCallback.java index ffbe09303..80cc6f8d7 100644 --- a/mcp/common/src/main/java/org/springframework/ai/mcp/SyncMcpToolCallback.java +++ b/mcp/common/src/main/java/org/springframework/ai/mcp/SyncMcpToolCallback.java @@ -23,6 +23,7 @@ import io.modelcontextprotocol.spec.McpSchema.CallToolRequest; import io.modelcontextprotocol.spec.McpSchema.CallToolResult; import io.modelcontextprotocol.spec.McpSchema.Tool; +import org.springframework.ai.chat.model.ToolContext; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.tool.ToolCallback; import org.springframework.ai.tool.definition.ToolDefinition; @@ -111,4 +112,10 @@ public class SyncMcpToolCallback implements ToolCallback { return ModelOptionsUtils.toJsonString(response.content()); } + @Override + public String call(String toolArguments, ToolContext toolContext) { + // ToolContext is not supported by the MCP tools + return this.call(toolArguments); + } + } diff --git a/mcp/common/src/test/java/org/springframework/ai/mcp/SyncMcpToolCallbackTests.java b/mcp/common/src/test/java/org/springframework/ai/mcp/SyncMcpToolCallbackTests.java index 713138525..70f2da83f 100644 --- a/mcp/common/src/test/java/org/springframework/ai/mcp/SyncMcpToolCallbackTests.java +++ b/mcp/common/src/test/java/org/springframework/ai/mcp/SyncMcpToolCallbackTests.java @@ -16,6 +16,8 @@ package org.springframework.ai.mcp; +import java.util.Map; + import io.modelcontextprotocol.client.McpSyncClient; import io.modelcontextprotocol.spec.McpSchema.CallToolRequest; import io.modelcontextprotocol.spec.McpSchema.CallToolResult; @@ -25,6 +27,8 @@ import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.ai.chat.model.ToolContext; + import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; @@ -71,4 +75,20 @@ class SyncMcpToolCallbackTests { assertThat(response).isNotNull(); } + @Test + void callShoulIngroeToolContext() { + // Arrange + when(tool.name()).thenReturn("testTool"); + CallToolResult callResult = mock(CallToolResult.class); + when(mcpClient.callTool(any(CallToolRequest.class))).thenReturn(callResult); + + SyncMcpToolCallback callback = new SyncMcpToolCallback(mcpClient, tool); + + // Act + String response = callback.call("{\"param\":\"value\"}", new ToolContext(Map.of("foo", "bar"))); + + // Assert + assertThat(response).isNotNull(); + } + }