diff --git a/contrib/spring-ai/src/main/java/com/google/adk/models/springai/ToolConverter.java b/contrib/spring-ai/src/main/java/com/google/adk/models/springai/ToolConverter.java index 4012ee5d6..6f4174a0e 100644 --- a/contrib/spring-ai/src/main/java/com/google/adk/models/springai/ToolConverter.java +++ b/contrib/spring-ai/src/main/java/com/google/adk/models/springai/ToolConverter.java @@ -16,6 +16,7 @@ package com.google.adk.models.springai; import com.google.adk.tools.BaseTool; +import com.google.adk.tools.ToolContext; import com.google.genai.types.FunctionDeclaration; import com.google.genai.types.Schema; import com.google.genai.types.Type; @@ -24,6 +25,7 @@ import java.util.List; import java.util.Map; import java.util.Set; +import java.util.function.BiFunction; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.tool.ToolCallback; @@ -40,6 +42,9 @@ public class ToolConverter { private static final Logger logger = LoggerFactory.getLogger(ToolConverter.class); + /** Key for passing an ADK {@link ToolContext} through Spring AI's tool context map. */ + public static final String ADK_TOOL_CONTEXT = "adk_tool_context"; + /** * Creates a tool registry from ADK tools for internal tracking. * @@ -120,35 +125,17 @@ public List convertToSpringAiTools(Map tools) { if (tool.declaration().isPresent()) { FunctionDeclaration declaration = tool.declaration().get(); - // Create a ToolCallback that wraps the ADK tool - // Create a Function that takes Map input and calls the ADK tool - java.util.function.Function, String> toolFunction = - args -> { - try { - logger.debug("Spring AI calling tool '{}'", tool.name()); - logger.debug("Raw args from Spring AI: {}", args); - logger.debug("Args type: {}", args.getClass().getName()); - logger.debug("Args keys: {}", args.keySet()); - for (Map.Entry entry : args.entrySet()) { - logger.debug( - " {} -> {} ({})", - entry.getKey(), - entry.getValue(), - entry.getValue().getClass().getName()); - } - - // Handle different argument formats that Spring AI might pass - Map processedArgs = processArguments(args, declaration); - logger.debug("Processed args for ADK: {}", processedArgs); - - // Call the ADK tool and wait for the result - Map result = tool.runAsync(processedArgs, null).blockingGet(); - // Convert result back to JSON string - return new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(result); - } catch (Exception e) { - throw new RuntimeException("Tool execution failed: " + e.getMessage(), e); - } - }; + BiFunction< + Map, + org.springframework.ai.chat.model.ToolContext, + Map> + toolFunction = + (args, springAiToolContext) -> { + logger.debug("Spring AI calling tool '{}'", tool.name()); + Map processedArgs = processArguments(args, declaration); + return tool.runAsync(processedArgs, adkToolContextFrom(springAiToolContext)) + .blockingGet(); + }; FunctionToolCallback.Builder callbackBuilder = FunctionToolCallback.builder(tool.name(), toolFunction).description(tool.description()); @@ -193,6 +180,26 @@ public List convertToSpringAiTools(Map tools) { return toolCallbacks; } + private ToolContext adkToolContextFrom( + org.springframework.ai.chat.model.ToolContext springAiToolContext) { + if (springAiToolContext == null) { + return null; + } + + Object context = springAiToolContext.getContext().get(ADK_TOOL_CONTEXT); + if (context == null) { + return null; + } + if (!(context instanceof ToolContext adkToolContext)) { + throw new IllegalArgumentException( + "Spring AI tool context entry '" + + ADK_TOOL_CONTEXT + + "' must be an ADK ToolContext, but was " + + context.getClass().getName()); + } + return adkToolContext; + } + /** * Process arguments from Spring AI format to ADK format. Spring AI might pass arguments in * different formats depending on the provider. diff --git a/contrib/spring-ai/src/test/java/com/google/adk/models/springai/ToolConverterTest.java b/contrib/spring-ai/src/test/java/com/google/adk/models/springai/ToolConverterTest.java index 1f3044159..116318e5f 100644 --- a/contrib/spring-ai/src/test/java/com/google/adk/models/springai/ToolConverterTest.java +++ b/contrib/spring-ai/src/test/java/com/google/adk/models/springai/ToolConverterTest.java @@ -16,17 +16,24 @@ package com.google.adk.models.springai; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.Mockito.mock; import com.google.adk.tools.BaseTool; +import com.google.adk.tools.ToolContext; import com.google.genai.types.FunctionDeclaration; import com.google.genai.types.Schema; +import io.reactivex.rxjava3.core.Single; import java.util.HashMap; import java.util.List; import java.util.Map; import java.util.Optional; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.springframework.ai.tool.ToolCallback; +import org.springframework.ai.tool.execution.ToolExecutionException; class ToolConverterTest { @@ -212,4 +219,103 @@ public Optional declaration() { assertThat(toolCallbacks).hasSize(1); assertThat(toolCallbacks.get(0).getToolDefinition().name()).isEqualTo("get_weather"); } + + @Test + void testToolCallbackWithoutAdkContextKeepsExistingNullContextBehavior() { + RecordingTool tool = new RecordingTool(); + ToolCallback callback = toolConverter.convertToSpringAiTools(Map.of(tool.name(), tool)).get(0); + + String directResult = callback.call("{\"location\":\"Paris\"}"); + String emptyContextResult = + callback.call( + "{\"location\":\"Paris\"}", + new org.springframework.ai.chat.model.ToolContext(Map.of())); + + assertThat(directResult).contains("sunny"); + assertThat(emptyContextResult).contains("sunny"); + assertThat(tool.invocationCount()).isEqualTo(2); + assertThat(tool.context()).isNull(); + assertThat(tool.arguments()).containsEntry("location", "Paris"); + } + + @Test + void testToolCallbackUsesAdkContextFromSpringAiToolContext() { + RecordingTool tool = new RecordingTool(); + ToolCallback callback = toolConverter.convertToSpringAiTools(Map.of(tool.name(), tool)).get(0); + ToolContext adkToolContext = mock(ToolContext.class); + org.springframework.ai.chat.model.ToolContext springAiToolContext = + new org.springframework.ai.chat.model.ToolContext( + Map.of(ToolConverter.ADK_TOOL_CONTEXT, adkToolContext)); + + String result = callback.call("{\"location\":\"Paris\"}", springAiToolContext); + + assertThat(result).contains("sunny"); + assertThat(tool.invocationCount()).isEqualTo(1); + assertThat(tool.context()).isSameAs(adkToolContext); + assertThat(tool.arguments()).containsEntry("location", "Paris"); + } + + @Test + void testToolCallbackRejectsWrongAdkContextType() { + RecordingTool tool = new RecordingTool(); + ToolCallback callback = toolConverter.convertToSpringAiTools(Map.of(tool.name(), tool)).get(0); + org.springframework.ai.chat.model.ToolContext springAiToolContext = + new org.springframework.ai.chat.model.ToolContext( + Map.of(ToolConverter.ADK_TOOL_CONTEXT, "not an ADK ToolContext")); + + assertThatThrownBy(() -> callback.call("{\"location\":\"Paris\"}", springAiToolContext)) + .isInstanceOf(ToolExecutionException.class) + .hasRootCauseInstanceOf(IllegalArgumentException.class) + .hasRootCauseMessage( + "Spring AI tool context entry 'adk_tool_context' must be an ADK ToolContext, but was java.lang.String"); + assertThat(tool.invocationCount()).isZero(); + } + + private static final class RecordingTool extends BaseTool { + private final FunctionDeclaration declaration; + private final AtomicInteger invocationCount = new AtomicInteger(); + private final AtomicReference> arguments = new AtomicReference<>(); + private final AtomicReference context = new AtomicReference<>(); + + private RecordingTool() { + super("get_weather", "Get weather for a location"); + this.declaration = + FunctionDeclaration.builder() + .name(name()) + .description(description()) + .parameters( + Schema.builder() + .type("OBJECT") + .properties(Map.of("location", Schema.builder().type("STRING").build())) + .required(List.of("location")) + .build()) + .build(); + } + + @Override + public Optional declaration() { + return Optional.of(declaration); + } + + @Override + public Single> runAsync( + Map arguments, ToolContext toolContext) { + invocationCount.incrementAndGet(); + this.arguments.set(arguments); + context.set(toolContext); + return Single.just(Map.of("forecast", "sunny")); + } + + private int invocationCount() { + return invocationCount.get(); + } + + private Map arguments() { + return arguments.get(); + } + + private ToolContext context() { + return context.get(); + } + } }