diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java index c79b611dc..2670b4a7f 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java @@ -53,6 +53,7 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.retry.RetryUtils; import org.springframework.http.ResponseEntity; import org.springframework.retry.support.RetryTemplate; @@ -413,8 +414,15 @@ public class AnthropicChatModel extends AbstractToolCallSupport implements ChatM systemPrompt, this.defaultOptions.getMaxTokens(), this.defaultOptions.getTemperature(), stream); if (prompt.getOptions() != null) { - AnthropicChatOptions updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), - ChatOptions.class, AnthropicChatOptions.class); + AnthropicChatOptions updatedRuntimeOptions; + if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions, + FunctionCallingOptions.class, AnthropicChatOptions.class); + } + else { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + AnthropicChatOptions.class); + } functionsForThisRequest.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions)); diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java index 62c4f2198..b96b56669 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java @@ -42,6 +42,7 @@ import org.springframework.ai.model.Media; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; import reactor.core.publisher.Flux; @@ -268,8 +269,15 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements Cha functionsForThisRequest.addAll(this.defaultOptions.getFunctions()); if (prompt.getOptions() != null) { - AzureOpenAiChatOptions updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), - ChatOptions.class, AzureOpenAiChatOptions.class); + AzureOpenAiChatOptions updatedRuntimeOptions; + if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions, + FunctionCallingOptions.class, AzureOpenAiChatOptions.class); + } + else { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + AzureOpenAiChatOptions.class); + } options = this.merge(updatedRuntimeOptions, options); functionsForThisRequest.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions)); diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java index 20d508629..e5d1cf72d 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java @@ -44,6 +44,7 @@ import org.springframework.ai.minimax.metadata.MiniMaxUsage; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.retry.RetryUtils; import org.springframework.http.ResponseEntity; import org.springframework.retry.support.RetryTemplate; @@ -391,8 +392,16 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod Set enabledToolsToUse = new HashSet<>(); if (prompt.getOptions() != null) { - MiniMaxChatOptions updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), - ChatOptions.class, MiniMaxChatOptions.class); + MiniMaxChatOptions updatedRuntimeOptions; + + if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions, + FunctionCallingOptions.class, MiniMaxChatOptions.class); + } + else { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + MiniMaxChatOptions.class); + } enabledToolsToUse.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions)); diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java index fc5b4170a..bad45cdd4 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java @@ -52,6 +52,7 @@ import org.springframework.ai.mistralai.metadata.MistralAiUsage; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.retry.RetryUtils; import org.springframework.http.ResponseEntity; import org.springframework.retry.support.RetryTemplate; @@ -367,8 +368,16 @@ public class MistralAiChatModel extends AbstractToolCallSupport implements ChatM request = ModelOptionsUtils.merge(request, this.defaultOptions, MistralAiApi.ChatCompletionRequest.class); if (prompt.getOptions() != null) { - var updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, - MistralAiChatOptions.class); + MistralAiChatOptions updatedRuntimeOptions; + + if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions, + FunctionCallingOptions.class, MistralAiChatOptions.class); + } + else { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + MistralAiChatOptions.class); + } functionsForThisRequest.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions)); diff --git a/models/spring-ai-moonshot/src/main/java/org/springframework/ai/moonshot/MoonshotChatModel.java b/models/spring-ai-moonshot/src/main/java/org/springframework/ai/moonshot/MoonshotChatModel.java index ce5ca15c7..ff1f22d75 100644 --- a/models/spring-ai-moonshot/src/main/java/org/springframework/ai/moonshot/MoonshotChatModel.java +++ b/models/spring-ai-moonshot/src/main/java/org/springframework/ai/moonshot/MoonshotChatModel.java @@ -33,6 +33,7 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.moonshot.api.MoonshotApi; import org.springframework.ai.moonshot.api.MoonshotApi.ChatCompletion; import org.springframework.ai.moonshot.api.MoonshotApi.ChatCompletion.Choice; @@ -341,9 +342,16 @@ public class MoonshotChatModel extends AbstractToolCallSupport implements ChatMo Set enabledToolsToUse = new HashSet<>(); if (prompt.getOptions() != null) { - MoonshotChatOptions updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), - ChatOptions.class, MoonshotChatOptions.class); + MoonshotChatOptions updatedRuntimeOptions; + if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions, + FunctionCallingOptions.class, MoonshotChatOptions.class); + } + else { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + MoonshotChatOptions.class); + } enabledToolsToUse.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions)); request = ModelOptionsUtils.merge(updatedRuntimeOptions, request, ChatCompletionRequest.class); diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java index 96e2a1267..c6d689e68 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java @@ -41,6 +41,7 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.ollama.api.OllamaApi; import org.springframework.ai.ollama.api.OllamaApi.ChatRequest; import org.springframework.ai.ollama.api.OllamaApi.Message.Role; @@ -297,8 +298,14 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode // runtime options OllamaOptions runtimeOptions = null; if (prompt.getOptions() != null) { - runtimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, - OllamaOptions.class); + if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) { + runtimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions, FunctionCallingOptions.class, + OllamaOptions.class); + } + else { + runtimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + OllamaOptions.class); + } functionsForThisRequest.addAll(this.runtimeFunctionCallbackConfigurations(runtimeOptions)); } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java index 5ac2f093f..340fcd79b 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java @@ -51,6 +51,7 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion.Choice; @@ -477,8 +478,16 @@ public class OpenAiChatModel extends AbstractToolCallSupport implements ChatMode Set enabledToolsToUse = new HashSet<>(); if (prompt.getOptions() != null) { - OpenAiChatOptions updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), - ChatOptions.class, OpenAiChatOptions.class); + OpenAiChatOptions updatedRuntimeOptions = null; + + if (prompt.getOptions() instanceof FunctionCallingOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(((FunctionCallingOptions) prompt.getOptions()), + FunctionCallingOptions.class, OpenAiChatOptions.class); + } + else if (prompt.getOptions() instanceof OpenAiChatOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + OpenAiChatOptions.class); + } enabledToolsToUse.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions)); diff --git a/models/spring-ai-qianfan/src/main/java/org/springframework/ai/qianfan/QianFanChatModel.java b/models/spring-ai-qianfan/src/main/java/org/springframework/ai/qianfan/QianFanChatModel.java index cfad456e4..54ec02670 100644 --- a/models/spring-ai-qianfan/src/main/java/org/springframework/ai/qianfan/QianFanChatModel.java +++ b/models/spring-ai-qianfan/src/main/java/org/springframework/ai/qianfan/QianFanChatModel.java @@ -181,15 +181,9 @@ public class QianFanChatModel implements ChatModel, StreamingChatModel { } if (prompt.getOptions() != null) { - if (prompt.getOptions() != null) { - var updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, - QianFanChatOptions.class); - request = ModelOptionsUtils.merge(updatedRuntimeOptions, request, ChatCompletionRequest.class); - } - else { - throw new IllegalArgumentException("Prompt options are not of type ChatOptions: " - + prompt.getOptions().getClass().getSimpleName()); - } + var updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + QianFanChatOptions.class); + request = ModelOptionsUtils.merge(updatedRuntimeOptions, request, ChatCompletionRequest.class); } return request; } diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java index 99051b463..b4a980e46 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java @@ -44,6 +44,7 @@ import org.springframework.ai.model.Media; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.retry.RetryUtils; import org.springframework.ai.vertexai.gemini.metadata.VertexAiUsage; import org.springframework.beans.factory.DisposableBean; @@ -71,8 +72,6 @@ import java.util.Set; */ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements ChatModel, DisposableBean { - private final static boolean IS_RUNTIME_CALL = true; - private final VertexAI vertexAI; private final VertexAiGeminiChatOptions defaultOptions; @@ -297,9 +296,15 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements VertexAiGeminiChatOptions updatedRuntimeOptions = VertexAiGeminiChatOptions.builder().build(); if (prompt.getOptions() != null) { - updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, - VertexAiGeminiChatOptions.class); + if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions, + FunctionCallingOptions.class, VertexAiGeminiChatOptions.class); + } + else { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + VertexAiGeminiChatOptions.class); + } functionsForThisRequest.addAll(runtimeFunctionCallbackConfigurations(updatedRuntimeOptions)); } diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java index c014f0506..2b2ed1ea0 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java @@ -33,6 +33,7 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; +import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.retry.RetryUtils; import org.springframework.ai.zhipuai.api.ZhiPuAiApi; import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletion; @@ -358,8 +359,15 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod Set enabledToolsToUse = new HashSet<>(); if (prompt.getOptions() != null) { - ZhiPuAiChatOptions updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), - ChatOptions.class, ZhiPuAiChatOptions.class); + ZhiPuAiChatOptions updatedRuntimeOptions; + if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions, + FunctionCallingOptions.class, ZhiPuAiChatOptions.class); + } + else { + updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class, + ZhiPuAiChatOptions.class); + } enabledToolsToUse.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions)); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java index f953e907d..347c699f9 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java @@ -18,10 +18,12 @@ package org.springframework.ai.model.function; import java.util.List; import java.util.Set; +import org.springframework.ai.chat.prompt.ChatOptions; + /** * @author Christian Tzolov */ -public interface FunctionCallingOptions { +public interface FunctionCallingOptions extends ChatOptions { /** * Function Callbacks to be registered with the ChatModel. For Prompt Options the diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithFunctionBeanIT.java index 284f04bf1..3a0a80052 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithFunctionBeanIT.java @@ -33,6 +33,7 @@ import org.springframework.ai.autoconfigure.anthropic.tool.MockWeatherService.Re import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; import org.springframework.context.annotation.Bean; @@ -79,6 +80,28 @@ class FunctionCallWithFunctionBeanIT { }); } + @Test + void functionCallWithPortableFunctionCallingOptions() { + + contextRunner + .withPropertyValues( + "spring.ai.anthropic.chat.options.model=" + AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue()) + .run(context -> { + + AnthropicChatModel chatModel = context.getBean(AnthropicChatModel.class); + + var userMessage = new UserMessage( + "What's the weather like in San Francisco, in Paris, France and in Tokyo, Japan? Return the temperature in Celsius."); + + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), + PortableFunctionCallingOptions.builder().withFunction("weatherFunction").build())); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); + }); + } + @Configuration static class Config { diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java index 799a37623..8b880fefd 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java @@ -30,6 +30,7 @@ import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; import org.springframework.context.annotation.Bean; @@ -80,6 +81,26 @@ class FunctionCallWithFunctionBeanIT { }); } + @Test + void functionCallWithPortableFunctionCallingOptions() { + contextRunner.withPropertyValues("spring.ai.azure.openai.chat.options..deployment-name=" + getDeploymentName()) + .run(context -> { + + ChatModel chatModel = context.getBean(AzureOpenAiChatModel.class); + + UserMessage userMessage = new UserMessage( + "What's the weather like in San Francisco, Paris and in Tokyo? Use Multi-turn function calling."); + + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), + PortableFunctionCallingOptions.builder().withFunction("weatherFunction").build())); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); + + }); + } + @Configuration static class Config { diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/WeatherServicePromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/WeatherServicePromptIT.java index 74e53163d..13750c3b0 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/WeatherServicePromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/WeatherServicePromptIT.java @@ -35,6 +35,8 @@ import org.springframework.ai.mistralai.MistralAiChatOptions; import org.springframework.ai.mistralai.api.MistralAiApi; import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionRequest.ToolChoice; import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.model.function.FunctionCallingOptions; +import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -63,7 +65,7 @@ public class WeatherServicePromptIT { MistralAiChatModel chatModel = context.getBean(MistralAiChatModel.class); - UserMessage userMessage = new UserMessage("What's the weather like in Paris?"); + UserMessage userMessage = new UserMessage("What's the weather like in Paris? Use Celsius."); // UserMessage userMessage = new UserMessage("What's the weather like in // San Francisco, Tokyo, and // Paris?"); @@ -86,6 +88,32 @@ public class WeatherServicePromptIT { }); } + @Test + void functionCallWithPortableFunctionCallingOptions() { + contextRunner + .withPropertyValues("spring.ai.mistralai.chat.options.model=" + MistralAiApi.ChatModel.LARGE.getValue()) + .run(context -> { + + MistralAiChatModel chatModel = context.getBean(MistralAiChatModel.class); + + UserMessage userMessage = new UserMessage("What's the weather like in Paris? Use Celsius."); + + PortableFunctionCallingOptions functionOptions = FunctionCallingOptions.builder() + .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MyWeatherService()) + .withName("CurrentWeatherService") + .withDescription("Get the current weather in requested location") + .build())) + + .build(); + + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions)); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).containsAnyOf("15", "15.0"); + }); + } + public static class MyWeatherService implements Function { // @formatter:off diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperIT.java index 0d42dfdbb..0064514e0 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/tool/FunctionCallbackWrapperIT.java @@ -35,6 +35,8 @@ import org.springframework.ai.chat.model.Generation; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.model.function.FunctionCallingOptions; +import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; import org.springframework.ai.ollama.OllamaChatModel; import org.springframework.ai.ollama.api.OllamaOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -125,6 +127,27 @@ public class FunctionCallbackWrapperIT { }); } + @Test + void functionCallWithPortableFunctionCallingOptions() { + contextRunner.run(context -> { + + OllamaChatModel chatModel = context.getBean(OllamaChatModel.class); + + // Test weatherFunction + UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + + PortableFunctionCallingOptions functionOptions = FunctionCallingOptions.builder() + .withFunction("WeatherInfo") + .build(); + + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions)); + + logger.info("Response: " + response.getResult().getOutput().getContent()); + + assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); + }); + } + @Configuration static class Config { diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java index ae29635fe..7ef80ca71 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java @@ -26,12 +26,13 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration; -import org.springframework.ai.chat.client.ChatClient; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.Generation; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.function.FunctionCallingOptions; +import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.OpenAiChatOptions; import org.springframework.ai.openai.api.OpenAiApi.ChatModel; @@ -91,15 +92,19 @@ class FunctionCallbackWithPlainFunctionBeanIT { OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); - // @formatter:off - String content = ChatClient.builder(chatModel).build().prompt() - .functions("weatherFunction") - .user("What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'weatherFunction'") - .stream().content() - .collectList().block().stream().collect(Collectors.joining()); - // @formatter:on + // Test weatherFunction + UserMessage userMessage = new UserMessage( + "What's the weather like in San Francisco, Tokyo, and Paris?"); - logger.info("Response: {}", content); + PortableFunctionCallingOptions functionOptions = FunctionCallingOptions.builder() + .withFunction("weatherFunction") + .build(); + + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions)); + + logger.info("Response: {}", response.getResult().getOutput().getContent()); + + assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java index 3657c8648..66a8c0113 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java @@ -28,6 +28,7 @@ import org.springframework.ai.autoconfigure.vertexai.gemini.VertexAiGeminiAutoCo import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel; import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -89,6 +90,39 @@ class FunctionCallWithFunctionBeanIT { }); } + @Test + void functionCallWithPortableFunctionCallingOptions() { + + contextRunner.withPropertyValues("spring.ai.vertex.ai.gemini.chat.options.model=" + // + VertexAiGeminiChatModel.ChatModel.GEMINI_PRO_1_5_PRO.getValue()) + + VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_FLASH.getValue()) + .run(context -> { + + VertexAiGeminiChatModel chatModel = context.getBean(VertexAiGeminiChatModel.class); + + var userMessage = new UserMessage(""" + What's the weather like in San Francisco, Paris and in Tokyo? + Return the temperature in Celsius. + Perform multiple funciton execution if necessary. + """); + + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), + PortableFunctionCallingOptions.builder().withFunction("weatherFunction").build())); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); + + response = chatModel.call(new Prompt(List.of(userMessage), + VertexAiGeminiChatOptions.builder().withFunction("weatherFunction3").build())); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); + + }); + } + @Configuration static class Config {