Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
Expand All @@ -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.
*
Expand Down Expand Up @@ -120,35 +125,17 @@ public List<ToolCallback> convertToSpringAiTools(Map<String, BaseTool> 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<Map<String, Object>, 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<String, Object> entry : args.entrySet()) {
logger.debug(
" {} -> {} ({})",
entry.getKey(),
entry.getValue(),
entry.getValue().getClass().getName());
}

// Handle different argument formats that Spring AI might pass
Map<String, Object> processedArgs = processArguments(args, declaration);
logger.debug("Processed args for ADK: {}", processedArgs);

// Call the ADK tool and wait for the result
Map<String, Object> 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<String, Object>,
org.springframework.ai.chat.model.ToolContext,
Map<String, Object>>
toolFunction =
(args, springAiToolContext) -> {
logger.debug("Spring AI calling tool '{}'", tool.name());
Map<String, Object> processedArgs = processArguments(args, declaration);
return tool.runAsync(processedArgs, adkToolContextFrom(springAiToolContext))
.blockingGet();
};

FunctionToolCallback.Builder callbackBuilder =
FunctionToolCallback.builder(tool.name(), toolFunction).description(tool.description());
Expand Down Expand Up @@ -193,6 +180,26 @@ public List<ToolCallback> convertToSpringAiTools(Map<String, BaseTool> 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(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This used to log and carry on; now one tool with an unserialisable schema fails every request, since convertToSpringAiTools runs on each one. Throwing is probably the better behaviour, but it's a behaviour change that isn't in the description and nothing tests it.

"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.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 {

Expand Down Expand Up @@ -212,4 +219,103 @@ public Optional<FunctionDeclaration> 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<Map<String, Object>> arguments = new AtomicReference<>();
private final AtomicReference<ToolContext> 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<FunctionDeclaration> declaration() {
return Optional.of(declaration);
}

@Override
public Single<Map<String, Object>> runAsync(
Map<String, Object> 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<String, Object> arguments() {
return arguments.get();
}

private ToolContext context() {
return context.get();
}
}
}