From 76c70189e2de7b8fb872ea0b01604258a2290e1f Mon Sep 17 00:00:00 2001 From: Sun Yuhan Date: Mon, 9 Jun 2025 10:48:16 +0800 Subject: [PATCH] fix: `validateToolContextSupport` method to use right order in ClassUtils.isAssignable Signed-off-by: Sun Yuhan --- .../ai/tool/method/MethodToolCallback.java | 2 +- .../MethodToolCallbackGenericTypesTest.java | 41 +++++++++++++++++++ 2 files changed, 42 insertions(+), 1 deletion(-) diff --git a/spring-ai-model/src/main/java/org/springframework/ai/tool/method/MethodToolCallback.java b/spring-ai-model/src/main/java/org/springframework/ai/tool/method/MethodToolCallback.java index cc320a54d..7c303f3a6 100644 --- a/spring-ai-model/src/main/java/org/springframework/ai/tool/method/MethodToolCallback.java +++ b/spring-ai-model/src/main/java/org/springframework/ai/tool/method/MethodToolCallback.java @@ -118,7 +118,7 @@ public final class MethodToolCallback implements ToolCallback { private void validateToolContextSupport(@Nullable ToolContext toolContext) { var isNonEmptyToolContextProvided = toolContext != null && !CollectionUtils.isEmpty(toolContext.getContext()); var isToolContextAcceptedByMethod = Stream.of(this.toolMethod.getParameterTypes()) - .anyMatch(type -> ClassUtils.isAssignable(type, ToolContext.class)); + .anyMatch(type -> ClassUtils.isAssignable(ToolContext.class, type)); if (isToolContextAcceptedByMethod && !isNonEmptyToolContextProvided) { throw new IllegalArgumentException("ToolContext is required by the method as an argument"); } diff --git a/spring-ai-model/src/test/java/org/springframework/ai/tool/method/MethodToolCallbackGenericTypesTest.java b/spring-ai-model/src/test/java/org/springframework/ai/tool/method/MethodToolCallbackGenericTypesTest.java index b99faa71a..6e05fd80c 100644 --- a/spring-ai-model/src/test/java/org/springframework/ai/tool/method/MethodToolCallbackGenericTypesTest.java +++ b/spring-ai-model/src/test/java/org/springframework/ai/tool/method/MethodToolCallbackGenericTypesTest.java @@ -22,6 +22,7 @@ import java.util.Map; import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.model.ToolContext; import org.springframework.ai.tool.definition.DefaultToolDefinition; import org.springframework.ai.tool.definition.ToolDefinition; @@ -137,6 +138,41 @@ class MethodToolCallbackGenericTypesTest { assertThat(result).isEqualTo("2 maps processed: [{a=1, b=2}, {c=3, d=4}]"); } + @Test + void testToolContextType() throws Exception { + // Create a test object with a method that takes a List> + TestGenericClass testObject = new TestGenericClass(); + Method method = TestGenericClass.class.getMethod("processStringListInToolContext", ToolContext.class); + + // Create a tool definition + ToolDefinition toolDefinition = DefaultToolDefinition.builder() + .name("processToolContext") + .description("Process tool context") + .inputSchema("{}") + .build(); + + // Create a MethodToolCallback + MethodToolCallback callback = MethodToolCallback.builder() + .toolDefinition(toolDefinition) + .toolMethod(method) + .toolObject(testObject) + .build(); + + // Create an empty JSON input + String toolInput = """ + {} + """; + + // Create a toolContext + ToolContext toolContext = new ToolContext(Map.of("foo", "bar")); + + // Call the tool + String result = callback.call(toolInput, toolContext); + + // Verify the result + assertThat(result).isEqualTo("1 entries processed {foo=bar}"); + } + /** * Test class with methods that use generic types. */ @@ -154,6 +190,11 @@ class MethodToolCallbackGenericTypesTest { return listOfMaps.size() + " maps processed: " + listOfMaps; } + public String processStringListInToolContext(ToolContext toolContext) { + Map context = toolContext.getContext(); + return context.size() + " entries processed " + context; + } + } }