feat: Improve validation in MethodToolCallbackProvider

This ensures that validation errors are caught early during object construction
rather than later when methods are called, providing better error feedback.

- Add validation for tool-annotated methods during construction
- Validate duplicate tool names in constructor instead of only at getToolCallbacks() time
- Add comprehensive test suite for MethodToolCallbackProvider

Signed-off-by: Christian Tzolov <christian.tzolov@broadcom.com>
This commit is contained in:
Christian Tzolov
2025-05-01 19:18:06 +03:00
committed by Mark Pollack
parent 90cab219d2
commit 6c52c99291
2 changed files with 161 additions and 0 deletions

View File

@@ -19,6 +19,7 @@ package org.springframework.ai.tool.method;
import java.lang.reflect.Method;
import java.util.Arrays;
import java.util.List;
import java.util.Optional;
import java.util.function.Consumer;
import java.util.function.Function;
import java.util.function.Supplier;
@@ -44,6 +45,7 @@ import org.springframework.util.ReflectionUtils;
* {@link Tool}-annotated methods.
*
* @author Thomas Vitale
* @author Christian Tzolov
* @since 1.0.0
*/
public final class MethodToolCallbackProvider implements ToolCallbackProvider {
@@ -55,7 +57,26 @@ public final class MethodToolCallbackProvider implements ToolCallbackProvider {
private MethodToolCallbackProvider(List<Object> toolObjects) {
Assert.notNull(toolObjects, "toolObjects cannot be null");
Assert.noNullElements(toolObjects, "toolObjects cannot contain null elements");
assertToolAnnotatedMethodsPresent(toolObjects);
this.toolObjects = toolObjects;
validateToolCallbacks(getToolCallbacks());
}
private void assertToolAnnotatedMethodsPresent(List<Object> toolObjects) {
for (Object toolObject : toolObjects) {
List<Method> toolMethods = Stream
.of(ReflectionUtils.getDeclaredMethods(
AopUtils.isAopProxy(toolObject) ? AopUtils.getTargetClass(toolObject) : toolObject.getClass()))
.filter(toolMethod -> toolMethod.isAnnotationPresent(Tool.class))
.filter(toolMethod -> !isFunctionalType(toolMethod))
.toList();
if (toolMethods.isEmpty()) {
throw new IllegalStateException("No @Tool annotated methods found in " + toolObject + "."
+ "Did you mean to pass a ToolCallback or ToolCallbackProvider? If so, you have to use .toolCallbacks() instead of .tool()");
}
}
}
@Override

View File

@@ -0,0 +1,140 @@
/*
* Copyright 2025-2025 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.tool.method;
import java.util.function.Consumer;
import java.util.function.Function;
import java.util.function.Supplier;
import org.junit.jupiter.api.Test;
import org.springframework.ai.tool.annotation.Tool;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
/**
* Unit tests for {@link MethodToolCallbackProvider}.
*
* @author Christian Tzolov
*/
class MethodToolCallbackProviderTests {
@Test
void whenToolObjectHasToolAnnotatedMethodThenSucceed() {
MethodToolCallbackProvider provider = MethodToolCallbackProvider.builder()
.toolObjects(new ValidToolObject())
.build();
assertThat(provider.getToolCallbacks()).hasSize(1);
assertThat(provider.getToolCallbacks()[0].getToolDefinition().name()).isEqualTo("validTool");
}
@Test
void whenToolObjectHasNoToolAnnotatedMethodThenThrow() {
assertThatThrownBy(
() -> MethodToolCallbackProvider.builder().toolObjects(new NoToolAnnotatedMethodObject()).build())
.isInstanceOf(IllegalStateException.class)
.hasMessageContaining("No @Tool annotated methods found in");
}
@Test
void whenToolObjectHasOnlyFunctionalTypeToolMethodsThenThrow() {
assertThatThrownBy(() -> MethodToolCallbackProvider.builder()
.toolObjects(new OnlyFunctionalTypeToolMethodsObject())
.build()).isInstanceOf(IllegalStateException.class)
.hasMessageContaining("No @Tool annotated methods found in");
}
@Test
void whenToolObjectHasMixOfValidAndFunctionalTypeToolMethodsThenSucceed() {
MethodToolCallbackProvider provider = MethodToolCallbackProvider.builder()
.toolObjects(new MixedToolMethodsObject())
.build();
assertThat(provider.getToolCallbacks()).hasSize(1);
assertThat(provider.getToolCallbacks()[0].getToolDefinition().name()).isEqualTo("validTool");
}
@Test
void whenMultipleToolObjectsWithSameToolNameThenThrow() {
assertThatThrownBy(() -> MethodToolCallbackProvider.builder()
.toolObjects(new ValidToolObject(), new DuplicateToolNameObject())
.build()).isInstanceOf(IllegalStateException.class)
.hasMessageContaining("Multiple tools with the same name (validTool) found in sources");
}
static class ValidToolObject {
@Tool
public String validTool() {
return "Valid tool result";
}
}
static class NoToolAnnotatedMethodObject {
public String notATool() {
return "Not a tool";
}
}
static class OnlyFunctionalTypeToolMethodsObject {
@Tool
public Function<String, String> functionTool() {
return input -> "Function result: " + input;
}
@Tool
public Supplier<String> supplierTool() {
return () -> "Supplier result";
}
@Tool
public Consumer<String> consumerTool() {
return input -> System.out.println("Consumer received: " + input);
}
}
static class MixedToolMethodsObject {
@Tool
public String validTool() {
return "Valid tool result";
}
@Tool
public Function<String, String> functionTool() {
return input -> "Function result: " + input;
}
}
static class DuplicateToolNameObject {
@Tool
public String validTool() {
return "Duplicate tool result";
}
}
}