diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/client/OpenAiChatClientMultipleFunctionCallsIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/client/OpenAiChatClientMultipleFunctionCallsIT.java index da045d0bb..63e55cc2a 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/client/OpenAiChatClientMultipleFunctionCallsIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/client/OpenAiChatClientMultipleFunctionCallsIT.java @@ -17,7 +17,9 @@ package org.springframework.ai.openai.chat.client; import static org.assertj.core.api.Assertions.assertThat; +import java.lang.reflect.Method; import java.util.List; +import java.util.function.Function; import java.util.stream.Collectors; import org.junit.jupiter.api.Test; @@ -25,6 +27,7 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.client.DefaultChatClient; import org.springframework.ai.openai.OpenAiTestConfiguration; import org.springframework.ai.openai.api.tool.MockWeatherService; import org.springframework.ai.openai.testutils.AbstractIT; @@ -123,4 +126,47 @@ class OpenAiChatClientMultipleFunctionCallsIT extends AbstractIT { } + @Test + void functionCallWithExplicitInputType() throws NoSuchMethodException { + + var chatClient = ChatClient.create(chatModel); + + Method currentTemp = MyFunction.class.getMethod("getCurrentTemp", MyFunction.Req.class); + + // NOTE: Lambda functions do not retain the type information, so we need to + // provide the input type explicitly. + MyFunction myFunction = new MyFunction(); + Function function = createFunction(myFunction, currentTemp); + + ChatClient.ChatClientRequestSpec chatClientRequestSpec = chatClient.prompt() + .user("What's the weather like in Shanghai?") + .function("currentTemp", "get current temp", MyFunction.Req.class, function); + + String content = chatClientRequestSpec.call().content(); + + assertThat(content).contains("23"); + } + + public static Function createFunction(Object obj, Method method) { + return (T t) -> { + try { + return (R) method.invoke(obj, t); + } + catch (Exception e) { + throw new RuntimeException(e); + } + }; + } + + public static class MyFunction { + + public record Req(String city) { + } + + public String getCurrentTemp(Req req) { + return "23"; + } + + } + } \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/client/ChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/client/ChatClient.java index ff826b3d5..a003d7039 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/client/ChatClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/client/ChatClient.java @@ -189,6 +189,9 @@ public interface ChatClient { ChatClientRequestSpec function(String name, String description, java.util.function.Function function); + ChatClientRequestSpec function(String name, String description, Class inputType, + java.util.function.Function function); + ChatClientRequestSpec functions(String... functionBeanNames); ChatClientRequestSpec system(String text); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/client/DefaultChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/client/DefaultChatClient.java index bca06eda7..63034db65 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/client/DefaultChatClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/client/DefaultChatClient.java @@ -610,6 +610,11 @@ public class DefaultChatClient implements ChatClient { public ChatClientRequestSpec function(String name, String description, java.util.function.Function function) { + return this.function(name, description, null, function); + } + + public ChatClientRequestSpec function(String name, String description, Class inputType, + java.util.function.Function function) { Assert.hasText(name, "the name must be non-null and non-empty"); Assert.hasText(description, "the description must be non-null and non-empty"); @@ -618,6 +623,7 @@ public class DefaultChatClient implements ChatClient { var fcw = FunctionCallbackWrapper.builder(function) .withDescription(description) .withName(name) + .withInputType(inputType) .withResponseConverter(Object::toString) .build(); this.functionCallbacks.add(fcw);