From 9db0ea7775a9ec99a10226294100f8d242bd44fb Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Tue, 21 May 2024 14:08:26 +0200 Subject: [PATCH] Add ChatCaller#getDefultOptions() breaking the funcitonal contract. - Implement fromOptions copty factory for eavery chatoptions implementation. - Extend ChatClient to use the caller default options if not provided explicitely. - Fix system/user text overdidign default system/user texts. Only non empty user/system text can overrid defult system/user text. - Rename chat() method to collect(). - Add OpenAI FunctionCallbackWrapper2IT auto-config tests. --- .../ai/anthropic/AnthropicChatOptions.java | 13 ++ .../ai/anthropic/AnthropicModelCaller.java | 5 + .../azure/openai/AzureOpenAiChatOptions.java | 18 ++ .../azure/openai/AzureOpenAiModelCaller.java | 2 +- .../anthropic/AnthropicChatOptions.java | 10 + .../BedrockAnthropicModelCaller.java | 5 + .../anthropic3/Anthropic3ChatOptions.java | 10 + .../BedrockAnthropic3ModelCaller.java | 6 + .../cohere/BedrockCohereChatOptions.java | 13 ++ .../cohere/BedrockCohereModelCaller.java | 5 + .../BedrockAi21Jurassic2ChatOptions.java | 15 ++ .../BedrockAi21Jurassic2ModelCaller.java | 5 + .../llama/BedrockLlamaChatOptions.java | 7 + .../llama/BedrockLlamaModelCaller.java | 5 + .../titan/BedrockTitanChatOptions.java | 8 + .../titan/BedrockTitanModelCaller.java | 5 + .../huggingface/HuggingfaceModelCaller.java | 7 + ...ChatClient.java => MiniMaxChatCaller.java} | 31 +-- .../ai/minimax/MiniMaxChatOptions.java | 18 ++ .../minimax/ChatCompletionRequestTests.java | 6 +- .../ai/minimax/MiniMaxTestConfiguration.java | 4 +- .../ai/minimax/api/MiniMaxRetryTests.java | 6 +- .../ai/mistralai/MistralAiChatOptions.java | 17 +- .../ai/mistralai/MistralAiModelCaller.java | 5 + .../ai/ollama/OllamaModelCaller.java | 5 + .../ai/ollama/api/OllamaOptions.java | 37 ++++ .../ai/openai/OpenAiChatOptions.java | 23 ++ .../ai/openai/OpenAiModelCaller.java | 5 + .../ai/openai/chat/OpenAiChatClientIT.java | 206 ++++++++++-------- .../gemini/VertexAiGeminiChatOptions.java | 14 ++ .../gemini/VertexAiGeminiModelCaller.java | 5 + .../palm2/VertexAiPaLm2ChatOptions.java | 9 + .../palm2/VertexAiPaLm2ModelCaller.java | 5 + .../ai/watsonx/WatsonxAiChatOptions.java | 16 ++ .../ai/watsonx/WatsonxAiModelCaller.java | 5 + ...ChatClient.java => ZhiPuAiChatCaller.java} | 15 +- .../ai/zhipuai/ZhiPuAiChatOptions.java | 20 ++ .../zhipuai/ChatCompletionRequestTests.java | 6 +- .../ai/zhipuai/ZhiPuAiTestConfiguration.java | 4 +- .../ai/zhipuai/api/ZhiPuAiRetryTests.java | 6 +- .../springframework/ai/chat/ChatCaller.java | 5 +- .../springframework/ai/chat/ChatClient.java | 92 ++++---- .../org/springframework/ai/chat/Main.java | 5 +- .../minimax/MiniMaxAutoConfiguration.java | 6 +- .../zhipuai/ZhiPuAiAutoConfiguration.java | 6 +- .../minimax/FunctionCallbackInPromptIT.java | 6 +- ...nctionCallbackWithPlainFunctionBeanIT.java | 8 +- .../minimax/FunctionCallbackWrapperIT.java | 6 +- .../minimax/MiniMaxAutoConfigurationIT.java | 6 +- .../minimax/MiniMaxPropertiesTests.java | 8 +- .../tool/FunctionCallbackWrapper2IT.java | 115 ++++++++++ .../tool/FunctionCallbackWrapperIT.java | 8 +- .../zhipuai/ZhiPuAiAutoConfigurationIT.java | 6 +- .../zhipuai/ZhiPuAiPropertiesTests.java | 8 +- .../tool/FunctionCallbackInPromptIT.java | 6 +- ...nctionCallbackWithPlainFunctionBeanIT.java | 8 +- .../tool/FunctionCallbackWrapperIT.java | 6 +- 57 files changed, 689 insertions(+), 223 deletions(-) rename models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/{MiniMaxChatClient.java => MiniMaxChatCaller.java} (93%) rename models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/{ZhiPuAiChatClient.java => ZhiPuAiChatCaller.java} (97%) create mode 100644 spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java index 025fae9c5..a79c92b69 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java @@ -223,4 +223,17 @@ public class AnthropicChatOptions implements ChatOptions, FunctionCallingOptions this.functions = functions; } + public static AnthropicChatOptions fromOptions(AnthropicChatOptions fromOptions) { + return builder().withModel(fromOptions.getModel()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withMetadata(fromOptions.getMetadata()) + .withStopSequences(fromOptions.getStopSequences()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicModelCaller.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicModelCaller.java index 67b095b6c..37deac66c 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicModelCaller.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicModelCaller.java @@ -457,4 +457,9 @@ public class AnthropicModelCaller extends "Streaming (stream=true) is not yet supported. We plan to add streaming support in a future beta version."); } + @Override + public ChatOptions getDefaultOptions() { + return AnthropicChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java index c477e3ce6..9dd7679c0 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java @@ -356,4 +356,22 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio this.functions = functions; } + public static AzureOpenAiChatOptions fromOptions(AzureOpenAiChatOptions fromOptions) { + return builder().withDeploymentName(fromOptions.getDeploymentName()) + .withFrequencyPenalty( + fromOptions.getFrequencyPenalty() != null ? fromOptions.getFrequencyPenalty().floatValue() : null) + .withLogitBias(fromOptions.getLogitBias()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withN(fromOptions.getN()) + .withPresencePenalty( + fromOptions.getPresencePenalty() != null ? fromOptions.getPresencePenalty().floatValue() : null) + .withStop(fromOptions.getStop()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withUser(fromOptions.getUser()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiModelCaller.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiModelCaller.java index 11c75b081..216c1eccb 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiModelCaller.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiModelCaller.java @@ -127,7 +127,7 @@ public class AzureOpenAiModelCaller } public AzureOpenAiChatOptions getDefaultOptions() { - return this.defaultOptions; + return AzureOpenAiChatOptions.fromOptions(this.defaultOptions); } @Override diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java index 2daceca21..d0d5a5a2c 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java @@ -164,4 +164,14 @@ public class AnthropicChatOptions implements ChatOptions { this.anthropicVersion = anthropicVersion; } + public static AnthropicChatOptions fromOptions(AnthropicChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withMaxTokensToSample(fromOptions.getMaxTokensToSample()) + .withTopK(fromOptions.getTopK()) + .withTopP(fromOptions.getTopP()) + .withStopSequences(fromOptions.getStopSequences()) + .withAnthropicVersion(fromOptions.getAnthropicVersion()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicModelCaller.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicModelCaller.java index 42b1a9594..78490468f 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicModelCaller.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicModelCaller.java @@ -117,4 +117,9 @@ public class BedrockAnthropicModelCaller implements ChatCaller, StreamingChatCal return request; } + @Override + public ChatOptions getDefaultOptions() { + return AnthropicChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java index 2862359fe..b4995683a 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java @@ -163,4 +163,14 @@ public class Anthropic3ChatOptions implements ChatOptions { this.anthropicVersion = anthropicVersion; } + public static Anthropic3ChatOptions fromOptions(Anthropic3ChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withTopK(fromOptions.getTopK()) + .withTopP(fromOptions.getTopP()) + .withStopSequences(fromOptions.getStopSequences()) + .withAnthropicVersion(fromOptions.getAnthropicVersion()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ModelCaller.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ModelCaller.java index fb593d529..7bdb5fc09 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ModelCaller.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ModelCaller.java @@ -15,6 +15,7 @@ */ package org.springframework.ai.bedrock.anthropic3; +import org.springframework.ai.bedrock.anthropic.AnthropicChatOptions; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.AnthropicChatRequest; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.AnthropicChatResponse; @@ -187,4 +188,9 @@ public class BedrockAnthropic3ModelCaller implements ChatCaller, StreamingChatCa } } + @Override + public ChatOptions getDefaultOptions() { + return Anthropic3ChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java index 89f625432..e0ab181cc 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java @@ -213,4 +213,17 @@ public class BedrockCohereChatOptions implements ChatOptions { this.truncate = truncate; } + public static BedrockCohereChatOptions fromOptions(BedrockCohereChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withStopSequences(fromOptions.getStopSequences()) + .withReturnLikelihoods(fromOptions.getReturnLikelihoods()) + .withNumGenerations(fromOptions.getNumGenerations()) + .withLogitBias(fromOptions.getLogitBias()) + .withTruncate(fromOptions.getTruncate()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereModelCaller.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereModelCaller.java index 7ac90b1c3..d11af0dcb 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereModelCaller.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereModelCaller.java @@ -114,4 +114,9 @@ public class BedrockCohereModelCaller implements ChatCaller, StreamingChatCaller return request; } + @Override + public ChatOptions getDefaultOptions() { + return BedrockCohereChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java index 4b62a5854..c165c61c1 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java @@ -413,4 +413,19 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions { } } + public static BedrockAi21Jurassic2ChatOptions fromOptions(BedrockAi21Jurassic2ChatOptions fromOptions) { + return builder().withPrompt(fromOptions.getPrompt()) + .withNumResults(fromOptions.getNumResults()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withMinTokens(fromOptions.getMinTokens()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .withStopSequences(fromOptions.getStopSequences()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withCountPenalty(fromOptions.getCountPenalty()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ModelCaller.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ModelCaller.java index e75f8616d..1ac7d3a8e 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ModelCaller.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ModelCaller.java @@ -121,4 +121,9 @@ public class BedrockAi21Jurassic2ModelCaller implements ChatCaller { } + @Override + public ChatOptions getDefaultOptions() { + return BedrockAi21Jurassic2ChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java index 3502fd4c4..4d6c0a6e0 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java @@ -109,4 +109,11 @@ public class BedrockLlamaChatOptions implements ChatOptions { throw new UnsupportedOperationException("Unsupported option: 'TopK'"); } + public static BedrockLlamaChatOptions fromOptions(BedrockLlamaChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withMaxGenLen(fromOptions.getMaxGenLen()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaModelCaller.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaModelCaller.java index 1a712e5dc..90a9b33f0 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaModelCaller.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaModelCaller.java @@ -130,4 +130,9 @@ public class BedrockLlamaModelCaller implements ChatCaller, StreamingChatCaller return request; } + @Override + public ChatOptions getDefaultOptions() { + return BedrockLlamaChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java index 61b16d16e..d53126a0b 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java @@ -128,4 +128,12 @@ public class BedrockTitanChatOptions implements ChatOptions { throw new UnsupportedOperationException("Bedrock Titan Chat does not support the 'TopK' option.'"); } + public static BedrockTitanChatOptions fromOptions(BedrockTitanChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withMaxTokenCount(fromOptions.getMaxTokenCount()) + .withStopSequences(fromOptions.getStopSequences()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanModelCaller.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanModelCaller.java index 3a7ea6937..0d2382925 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanModelCaller.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanModelCaller.java @@ -146,4 +146,9 @@ public class BedrockTitanModelCaller implements ChatCaller, StreamingChatCaller }; } + @Override + public ChatOptions getDefaultOptions() { + return BedrockTitanChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceModelCaller.java b/models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceModelCaller.java index d0c603e9f..3b8dfc184 100644 --- a/models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceModelCaller.java +++ b/models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceModelCaller.java @@ -31,6 +31,8 @@ import org.springframework.ai.huggingface.model.AllOfGenerateResponseDetails; import org.springframework.ai.huggingface.model.GenerateParameters; import org.springframework.ai.huggingface.model.GenerateRequest; import org.springframework.ai.huggingface.model.GenerateResponse; +import org.springframework.ai.chat.prompt.ChatOptions; +import org.springframework.ai.chat.prompt.ChatOptionsBuilder; import org.springframework.ai.chat.prompt.Prompt; /** @@ -120,4 +122,9 @@ public class HuggingfaceModelCaller implements ChatCaller { this.maxNewTokens = maxNewTokens; } + @Override + public ChatOptions getDefaultOptions() { + return ChatOptionsBuilder.builder().build(); + } + } diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatClient.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatCaller.java similarity index 93% rename from models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatClient.java rename to models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatCaller.java index aae3f29af..a9049d500 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatClient.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatCaller.java @@ -58,11 +58,11 @@ import java.util.concurrent.ConcurrentHashMap; * @see MiniMaxApi * @since 1.0.0 M1 */ -public class MiniMaxChatClient extends +public class MiniMaxChatCaller extends AbstractFunctionCallSupport> implements ChatCaller, StreamingChatCaller { - private static final Logger logger = LoggerFactory.getLogger(MiniMaxChatClient.class); + private static final Logger logger = LoggerFactory.getLogger(MiniMaxChatCaller.class); /** * The default options used for the chat completion requests. @@ -85,7 +85,7 @@ public class MiniMaxChatClient extends * MiniMax Chat API. * @throws IllegalArgumentException if MiniMaxApi is null */ - public MiniMaxChatClient(MiniMaxApi miniMaxApi) { + public MiniMaxChatCaller(MiniMaxApi miniMaxApi) { this(miniMaxApi, MiniMaxChatOptions.builder().withModel(MiniMaxApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); } @@ -96,7 +96,7 @@ public class MiniMaxChatClient extends * MiniMax Chat API. * @param options The MiniMaxChatOptions to configure the chat client. */ - public MiniMaxChatClient(MiniMaxApi miniMaxApi, MiniMaxChatOptions options) { + public MiniMaxChatCaller(MiniMaxApi miniMaxApi, MiniMaxChatOptions options) { this(miniMaxApi, options, null, RetryUtils.DEFAULT_RETRY_TEMPLATE); } @@ -108,7 +108,7 @@ public class MiniMaxChatClient extends * @param functionCallbackContext The function callback context. * @param retryTemplate The retry template. */ - public MiniMaxChatClient(MiniMaxApi miniMaxApi, MiniMaxChatOptions options, + public MiniMaxChatCaller(MiniMaxApi miniMaxApi, MiniMaxChatOptions options, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { super(functionCallbackContext); Assert.notNull(miniMaxApi, "MiniMaxApi must not be null"); @@ -281,22 +281,6 @@ public class MiniMaxChatClient extends return request; } - private String fromMediaData(MimeType mimeType, Object mediaContentData) { - if (mediaContentData instanceof byte[] bytes) { - // Assume the bytes are an image. So, convert the bytes to a base64 encoded - // following the prefix pattern. - return String.format("data:%s;base64,%s", mimeType.toString(), Base64.getEncoder().encodeToString(bytes)); - } - else if (mediaContentData instanceof String text) { - // Assume the text is a URLs or a base64 encoded image prefixed by the user. - return text; - } - else { - throw new IllegalArgumentException( - "Unsupported media data type: " + mediaContentData.getClass().getSimpleName()); - } - } - private List getFunctionTools(Set functionNames) { return this.resolveFunctionCallbacks(functionNames).stream().map(functionCallback -> { var function = new MiniMaxApi.FunctionTool.Function(functionCallback.getDescription(), @@ -376,4 +360,9 @@ public class MiniMaxChatClient extends && choice.finishReason() == MiniMaxApi.ChatCompletionFinishReason.TOOL_CALLS; } + @Override + public ChatOptions getDefaultOptions() { + return MiniMaxChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java index 1b46ca6a6..d6bb27c9b 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java @@ -467,4 +467,22 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions { return true; } + public static MiniMaxChatOptions fromOptions(MiniMaxChatOptions fromOptions) { + return builder().withModel(fromOptions.getModel()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withN(fromOptions.getN()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withResponseFormat(fromOptions.getResponseFormat()) + .withSeed(fromOptions.getSeed()) + .withStop(fromOptions.getStop()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTools(fromOptions.getTools()) + .withToolChoice(fromOptions.getToolChoice()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/ChatCompletionRequestTests.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/ChatCompletionRequestTests.java index 9adf803a4..3dd8c9fc1 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/ChatCompletionRequestTests.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/ChatCompletionRequestTests.java @@ -33,7 +33,7 @@ public class ChatCompletionRequestTests { @Test public void createRequestWithChatOptions() { - var client = new MiniMaxChatClient(new MiniMaxApi("TEST"), + var client = new MiniMaxChatCaller(new MiniMaxApi("TEST"), MiniMaxChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); var request = client.createRequest(new Prompt("Test message content"), false); @@ -59,7 +59,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new MiniMaxChatClient(new MiniMaxApi("TEST"), + var client = new MiniMaxChatCaller(new MiniMaxApi("TEST"), MiniMaxChatOptions.builder().withModel("DEFAULT_MODEL").build()); var request = client.createRequest(new Prompt("Test message content", @@ -89,7 +89,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new MiniMaxChatClient(new MiniMaxApi("TEST"), + var client = new MiniMaxChatCaller(new MiniMaxApi("TEST"), MiniMaxChatOptions.builder() .withModel("DEFAULT_MODEL") .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/MiniMaxTestConfiguration.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/MiniMaxTestConfiguration.java index f544b4896..3d3efdd95 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/MiniMaxTestConfiguration.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/MiniMaxTestConfiguration.java @@ -42,8 +42,8 @@ public class MiniMaxTestConfiguration { } @Bean - public MiniMaxChatClient miniMaxChatClient(MiniMaxApi api) { - return new MiniMaxChatClient(api); + public MiniMaxChatCaller miniMaxChatClient(MiniMaxApi api) { + return new MiniMaxChatCaller(api); } @Bean diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java index 46d62ae3c..2a5f31c52 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java @@ -22,7 +22,7 @@ import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatCaller; import org.springframework.ai.minimax.MiniMaxChatOptions; import org.springframework.ai.minimax.MiniMaxEmbeddingClient; import org.springframework.ai.minimax.MiniMaxEmbeddingOptions; @@ -83,7 +83,7 @@ public class MiniMaxRetryTests { private @Mock MiniMaxApi miniMaxApi; - private MiniMaxChatClient chatClient; + private MiniMaxChatCaller chatClient; private MiniMaxEmbeddingClient embeddingClient; @@ -93,7 +93,7 @@ public class MiniMaxRetryTests { retryListener = new TestRetryListener(); retryTemplate.registerListener(retryListener); - chatClient = new MiniMaxChatClient(miniMaxApi, MiniMaxChatOptions.builder().build(), null, retryTemplate); + chatClient = new MiniMaxChatCaller(miniMaxApi, MiniMaxChatOptions.builder().build(), null, retryTemplate); embeddingClient = new MiniMaxEmbeddingClient(miniMaxApi, MetadataMode.EMBED, MiniMaxEmbeddingOptions.builder().build(), retryTemplate); } diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java index d16f3daee..1b3034e38 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java @@ -139,7 +139,7 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions return this; } - public Builder withMaxToken(Integer maxTokens) { + public Builder withMaxTokens(Integer maxTokens) { this.options.setMaxTokens(maxTokens); return this; } @@ -309,4 +309,19 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions this.functions = functions; } + public static MistralAiChatOptions fromOptions(MistralAiChatOptions fromOptions) { + return builder().withModel(fromOptions.getModel()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withSafePrompt(fromOptions.getSafePrompt()) + .withRandomSeed(fromOptions.getRandomSeed()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withResponseFormat(fromOptions.getResponseFormat()) + .withTools(fromOptions.getTools()) + .withToolChoice(fromOptions.getToolChoice()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiModelCaller.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiModelCaller.java index 3f238b0ad..5853ae174 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiModelCaller.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiModelCaller.java @@ -324,4 +324,9 @@ public class MistralAiModelCaller extends return !CollectionUtils.isEmpty(choices.get(0).message().toolCalls()); } + @Override + public ChatOptions getDefaultOptions() { + return MistralAiChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaModelCaller.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaModelCaller.java index 2b60725b5..36c627b30 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaModelCaller.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaModelCaller.java @@ -205,4 +205,9 @@ public class OllamaModelCaller implements ChatCaller, StreamingChatCaller { } } + @Override + public ChatOptions getDefaultOptions() { + return OllamaOptions.fromOptions(this.defaultOptions); + } + } \ No newline at end of file diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java index abc329e81..436631f6e 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java @@ -714,6 +714,43 @@ public class OllamaOptions implements ChatOptions, EmbeddingOptions { .collect(Collectors.toMap(Map.Entry::getKey, Map.Entry::getValue)); } + public static OllamaOptions fromOptions(OllamaOptions fromOptions) { + return new OllamaOptions() + .withModel(fromOptions.getModel()) + .withFormat(fromOptions.getFormat()) + .withKeepAlive(fromOptions.getKeepAlive()) + .withUseNUMA(fromOptions.getUseNUMA()) + .withNumCtx(fromOptions.getNumCtx()) + .withNumBatch(fromOptions.getNumBatch()) + .withNumGQA(fromOptions.getNumGQA()) + .withNumGPU(fromOptions.getNumGPU()) + .withMainGPU(fromOptions.getMainGPU()) + .withLowVRAM(fromOptions.getLowVRAM()) + .withF16KV(fromOptions.getF16KV()) + .withLogitsAll(fromOptions.getLogitsAll()) + .withVocabOnly(fromOptions.getVocabOnly()) + .withUseMMap(fromOptions.getUseMMap()) + .withUseMLock(fromOptions.getUseMLock()) + .withNumThread(fromOptions.getNumThread()) + .withNumKeep(fromOptions.getNumKeep()) + .withSeed(fromOptions.getSeed()) + .withNumPredict(fromOptions.getNumPredict()) + .withTopK(fromOptions.getTopK()) + .withTopP(fromOptions.getTopP()) + .withTfsZ(fromOptions.getTfsZ()) + .withTypicalP(fromOptions.getTypicalP()) + .withRepeatLastN(fromOptions.getRepeatLastN()) + .withTemperature(fromOptions.getTemperature()) + .withRepeatPenalty(fromOptions.getRepeatPenalty()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withMirostat(fromOptions.getMirostat()) + .withMirostatTau(fromOptions.getMirostatTau()) + .withMirostatEta(fromOptions.getMirostatEta()) + .withPenalizeNewline(fromOptions.getPenalizeNewline()) + .withStop(fromOptions.getStop()); + } + // @formatter:on diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java index 8447f7410..e50aaa711 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java @@ -567,4 +567,27 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions { throw new UnsupportedOperationException("Unimplemented method 'setTopK'"); } + public static OpenAiChatOptions fromOptions(OpenAiChatOptions fromOptions) { + return OpenAiChatOptions.builder() + .withModel(fromOptions.getModel()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withLogitBias(fromOptions.getLogitBias()) + .withLogprobs(fromOptions.getLogprobs()) + .withTopLogprobs(fromOptions.getTopLogprobs()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withN(fromOptions.getN()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withResponseFormat(fromOptions.getResponseFormat()) + .withSeed(fromOptions.getSeed()) + .withStop(fromOptions.getStop()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTools(fromOptions.getTools()) + .withToolChoice(fromOptions.getToolChoice()) + .withUser(fromOptions.getUser()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiModelCaller.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiModelCaller.java index 7102a2abb..3f8385a89 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiModelCaller.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiModelCaller.java @@ -394,4 +394,9 @@ public class OpenAiModelCaller extends && choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; } + @Override + public ChatOptions getDefaultOptions() { + return OpenAiChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java index f5adb6631..9a0e9c98e 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java @@ -62,10 +62,13 @@ class OpenAiChatClientIT extends AbstractIT { @Test void roleTest() { - ChatResponse response = ChatClient.builder(modelCaller).build().call() - .system(s -> s.text(systemTextResource).param("name", "Bob").param("voice", "pirate")) - .user(u -> u.text("Tell me about 3 famous pirates from the Golden Age of Piracy and what they did")) - .chat().chatResponse(); + ChatResponse response = ChatClient.builder(modelCaller) + .build() + .call() + .system(s -> s.text(systemTextResource).param("name", "Bob").param("voice", "pirate")) + .user(u -> u.text("Tell me about 3 famous pirates from the Golden Age of Piracy and what they did")) + .collect() + .chatResponse(); System.out.println(response); // UserMessage userMessage = new UserMessage( @@ -86,9 +89,12 @@ class OpenAiChatClientIT extends AbstractIT { void listOutputConverter() { // TODO: there is a problem here. - Collection list = ChatClient.builder(modelCaller).build().call() - .user(u -> u.text("List five {subject}").param("subject", "ice cream flavors")) - .chat().list(String.class); + Collection list = ChatClient.builder(modelCaller) + .build() + .call() + .user(u -> u.text("List five {subject}").param("subject", "ice cream flavors")) + .collect() + .list(String.class); // DefaultConversionService conversionService = new DefaultConversionService(); // ListOutputConverter outputConverter = new @@ -113,11 +119,14 @@ class OpenAiChatClientIT extends AbstractIT { @Test void mapOutputConverter() { - Map result = ChatClient.builder(modelCaller).build().call() - .user(u -> u.text("Provide me a List of {subject}") - .param("subject", "an array of numbers from 1 to 9 under they key name 'numbers'")) - .chat().single(new ParameterizedTypeReference>() { - }); + Map result = ChatClient.builder(modelCaller) + .build() + .call() + .user(u -> u.text("Provide me a List of {subject}") + .param("subject", "an array of numbers from 1 to 9 under they key name 'numbers'")) + .collect() + .single(new ParameterizedTypeReference>() { + }); // MapOutputConverter outputConverter = new MapOutputConverter(); @@ -140,9 +149,12 @@ class OpenAiChatClientIT extends AbstractIT { @Test void beanOutputConverter() { - ActorsFilms actorsFilms = ChatClient.builder(modelCaller).build().call() - .user(u -> u.text("Generate the filmography for a random actor.")) - .chat().single(ActorsFilms.class); + ActorsFilms actorsFilms = ChatClient.builder(modelCaller) + .build() + .call() + .user(u -> u.text("Generate the filmography for a random actor.")) + .collect() + .single(ActorsFilms.class); // BeanOutputConverter outputConverter = new // BeanOutputConverter<>(ActorsFilms.class); @@ -170,9 +182,11 @@ class OpenAiChatClientIT extends AbstractIT { void beanOutputConverterRecords() { ActorsFilmsRecord actorsFilms = ChatClient.builder(modelCaller) - .build().call() - .user(u -> u.text("Generate the filmography of 5 movies for Tom Hanks.")) - .chat().single(ActorsFilmsRecord.class); + .build() + .call() + .user(u -> u.text("Generate the filmography of 5 movies for Tom Hanks.")) + .collect() + .single(ActorsFilmsRecord.class); // BeanOutputConverter outputConverter = new // BeanOutputConverter<>(ActorsFilmsRecord.class); @@ -199,31 +213,38 @@ class OpenAiChatClientIT extends AbstractIT { BeanOutputConverter outputConverter = new BeanOutputConverter<>(ActorsFilmsRecord.class); - Flux chatResponse = ChatClient.builder(modelCaller).build() - .call().user(u -> u.text("Generate the filmography of 5 movies for Tom Hanks. " + System.lineSeparator() + "{format}") - .param("format", outputConverter.getFormat())) - .stream().chatResponse(); + Flux chatResponse = ChatClient.builder(modelCaller) + .build() + .call() + .user(u -> u + .text("Generate the filmography of 5 movies for Tom Hanks. " + System.lineSeparator() + "{format}") + .param("format", outputConverter.getFormat())) + .stream() + .chatResponse(); - String generationTextFromStream = chatResponse.collectList().block() - .stream() - .map(ChatResponse::getResult) - .map(Generation::getOutput) - .map(AssistantMessage::getContent) - .collect(Collectors.joining()); + String generationTextFromStream = chatResponse.collectList() + .block() + .stream() + .map(ChatResponse::getResult) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); // String generationTextFromStream = chatResponse.collectList() // .block() // .stream() // .collect(Collectors.joining()); - // BeanOutputConverter outputConverter = new BeanOutputConverter<>(ActorsFilmsRecord.class); + // BeanOutputConverter outputConverter = new + // BeanOutputConverter<>(ActorsFilmsRecord.class); // String format = outputConverter.getFormat(); // String template = """ // Generate the filmography of 5 movies for Tom Hanks. // {format} // """; - // PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); + // PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", + // format)); // Prompt prompt = new Prompt(promptTemplate.createMessage()); // String generationTextFromStream = streamingChatClient.stream(prompt) @@ -245,13 +266,13 @@ class OpenAiChatClientIT extends AbstractIT { @Test void functionCallTest() { - ChatResponse response = ChatClient.builder(modelCaller).build().call() - .user(u -> u.text("What's the weather like in San Francisco, Tokyo, and Paris?")) - // TODO how to use the protable function calling options internally. - // Perhaps the ModelCaller a emptyOptions() method needs to be provided. - .options(OpenAiChatOptions.builder().build()) - .function("getCurrentWeather", "Get the weather in location", new MockWeatherService()) - .chat().chatResponse(); + ChatResponse response = ChatClient.builder(modelCaller) + .build() + .call() + .user(u -> u.text("What's the weather like in San Francisco, Tokyo, and Paris?")) + .function("getCurrentWeather", "Get the weather in location", new MockWeatherService()) + .collect() + .chatResponse(); // UserMessage userMessage = new UserMessage("What's the weather like in San // Francisco, Tokyo, and Paris?"); @@ -280,37 +301,40 @@ class OpenAiChatClientIT extends AbstractIT { @Test void streamFunctionCallTest() { - Flux response = ChatClient.builder(modelCaller).build().call() - .user(u -> u.text("What's the weather like in San Francisco, Tokyo, and Paris?")) - // TODO how to use the protable function calling options internally. - // Perhaps the ModelCaller a emptyOptions() method needs to be provided. - .options(OpenAiChatOptions.builder().build()) - .function("getCurrentWeather", "Get the weather in location", new MockWeatherService()) - .stream().chatResponse(); + Flux response = ChatClient.builder(modelCaller) + .build() + .call() + .user(u -> u.text("What's the weather like in San Francisco, Tokyo, and Paris?")) + .function("getCurrentWeather", "Get the weather in location", new MockWeatherService()) + .stream() + .chatResponse(); - // UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + // UserMessage userMessage = new UserMessage("What's the weather like in San + // Francisco, Tokyo, and Paris?"); // List messages = new ArrayList<>(List.of(userMessage)); // var promptOptions = OpenAiChatOptions.builder() // // .withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue()) - // .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) + // .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new + // MockWeatherService()) // .withName("getCurrentWeather") // .withDescription("Get the weather in location") // .withResponseConverter((response) -> "" + response.temp() + response.unit()) // .build())) // .build(); - // Flux response = streamingChatClient.stream(new Prompt(messages, promptOptions)); + // Flux response = streamingChatClient.stream(new Prompt(messages, + // promptOptions)); String content = response.collectList() - .block() - .stream() - .map(ChatResponse::getResults) - .flatMap(List::stream) - .map(Generation::getOutput) - .map(AssistantMessage::getContent) - .collect(Collectors.joining()); + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); logger.info("Response: {}", content); assertThat(content).containsAnyOf("30.0", "30"); @@ -323,15 +347,15 @@ class OpenAiChatClientIT extends AbstractIT { void multiModalityEmbeddedImage(String modelName) throws IOException { ChatResponse response = ChatClient.builder(modelCaller) - .build() - .call() - // TODO consider adding model(...) method to ChatClient as a shortcut to - // OpenAiChatOptions.builder().withModel(modelName).build() - .options(OpenAiChatOptions.builder().withModel(modelName).build()) - .user(u -> u.text("Explain what do you see on this picture?") - .media(MimeTypeUtils.IMAGE_PNG, new ClassPathResource("/test.png"))) - .chat() - .chatResponse(); + .build() + .call() + // TODO consider adding model(...) method to ChatClient as a shortcut to + // OpenAiChatOptions.builder().withModel(modelName).build() + .options(OpenAiChatOptions.builder().withModel(modelName).build()) + .user(u -> u.text("Explain what do you see on this picture?") + .media(MimeTypeUtils.IMAGE_PNG, new ClassPathResource("/test.png"))) + .collect() + .chatResponse(); // var imageData = new ClassPathResource("/test.png"); @@ -354,13 +378,15 @@ class OpenAiChatClientIT extends AbstractIT { // TODO: add url method that wrapps the checked exception. URL url = new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"); - ChatResponse response = ChatClient.builder(modelCaller).build().call() - // TODO consider adding model(...) method to ChatClient as a shortcut to - // OpenAiChatOptions.builder().withModel(modelName).build() - .options(OpenAiChatOptions.builder().withModel(modelName).build()) - .user(u -> u.text("Explain what do you see on this picture?") - .media(MimeTypeUtils.IMAGE_PNG, url)) - .chat().chatResponse(); + ChatResponse response = ChatClient.builder(modelCaller) + .build() + .call() + // TODO consider adding model(...) method to ChatClient as a shortcut to + // OpenAiChatOptions.builder().withModel(modelName).build() + .options(OpenAiChatOptions.builder().withModel(modelName).build()) + .user(u -> u.text("Explain what do you see on this picture?").media(MimeTypeUtils.IMAGE_PNG, url)) + .collect() + .chatResponse(); // var userMessage = new UserMessage("Explain what do you see on this picture?", // List @@ -383,30 +409,34 @@ class OpenAiChatClientIT extends AbstractIT { // TODO: add url method that wrapps the checked exception. URL url = new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"); - Flux response = ChatClient.builder(modelCaller).build().call() - // TODO consider adding model(...) method to ChatClient as a shortcut to - // OpenAiChatOptions.builder().withModel(modelName).build() - .options(OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_VISION_PREVIEW.getValue()) - .build()) - .user(u -> u.text("Explain what do you see on this picture?") - .media(MimeTypeUtils.IMAGE_PNG, url)) - .stream().chatResponse(); + Flux response = ChatClient.builder(modelCaller) + .build() + .call() + // TODO consider adding model(...) method to ChatClient as a shortcut to + // OpenAiChatOptions.builder().withModel(modelName).build() + .options(OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_VISION_PREVIEW.getValue()).build()) + .user(u -> u.text("Explain what do you see on this picture?").media(MimeTypeUtils.IMAGE_PNG, url)) + .stream() + .chatResponse(); - // var userMessage = new UserMessage("Explain what do you see on this picture?", List + // var userMessage = new UserMessage("Explain what do you see on this picture?", + // List // .of(new Media(MimeTypeUtils.IMAGE_PNG, - // new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png")))); + // new + // URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png")))); - // Flux response = streamingChatClient.stream(new Prompt(List.of(userMessage), + // Flux response = streamingChatClient.stream(new + // Prompt(List.of(userMessage), // OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_VISION_PREVIEW.getValue()).build())); String content = response.collectList() - .block() - .stream() - .map(ChatResponse::getResults) - .flatMap(List::stream) - .map(Generation::getOutput) - .map(AssistantMessage::getContent) - .collect(Collectors.joining()); + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); logger.info("Response: {}", content); assertThat(content).contains("bananas", "apple"); assertThat(content).containsAnyOf("bowl", "basket"); diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java index 081136c42..dfe45a27f 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java @@ -336,4 +336,18 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp return true; } + public static VertexAiGeminiChatOptions fromOptions(VertexAiGeminiChatOptions fromOptions) { + VertexAiGeminiChatOptions options = new VertexAiGeminiChatOptions(); + options.setStopSequences(fromOptions.getStopSequences()); + options.setTemperature(fromOptions.getTemperature()); + options.setTopP(fromOptions.getTopP()); + options.setTopK(fromOptions.getTopK()); + options.setCandidateCount(fromOptions.getCandidateCount()); + options.setMaxOutputTokens(fromOptions.getMaxOutputTokens()); + options.setModel(fromOptions.getModel()); + options.setFunctionCallbacks(fromOptions.getFunctionCallbacks()); + options.setFunctions(fromOptions.getFunctions()); + return options; + } + } diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiModelCaller.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiModelCaller.java index 3292656f6..7000032f8 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiModelCaller.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiModelCaller.java @@ -478,4 +478,9 @@ public class VertexAiGeminiModelCaller return response.getCandidatesList().get(0).getContent().getPartsList().get(0).hasFunctionCall(); } + @Override + public ChatOptions getDefaultOptions() { + return VertexAiGeminiChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java index cae09e7bb..55742ff05 100644 --- a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java +++ b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java @@ -127,4 +127,13 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions { this.topK = topK; } + public static VertexAiPaLm2ChatOptions fromOptions(VertexAiPaLm2ChatOptions fromOptions) { + return VertexAiPaLm2ChatOptions.builder() + .withTemperature(fromOptions.getTemperature()) + .withCandidateCount(fromOptions.getCandidateCount()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .build(); + } + } diff --git a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ModelCaller.java b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ModelCaller.java index e047d7323..0ae265932 100644 --- a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ModelCaller.java +++ b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ModelCaller.java @@ -111,4 +111,9 @@ public class VertexAiPaLm2ModelCaller implements ChatCaller { return request; } + @Override + public ChatOptions getDefaultOptions() { + return VertexAiPaLm2ChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java index 0b10febdb..28bd02159 100644 --- a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java +++ b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java @@ -324,5 +324,21 @@ public class WatsonxAiChatOptions implements ChatOptions { return input != null ? input.replaceAll("([a-z])([A-Z]+)", "$1_$2").toLowerCase() : null; } + public static WatsonxAiChatOptions fromOptions(WatsonxAiChatOptions fromOptions) { + return WatsonxAiChatOptions.builder() + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .withDecodingMethod(fromOptions.getDecodingMethod()) + .withMaxNewTokens(fromOptions.getMaxNewTokens()) + .withMinNewTokens(fromOptions.getMinNewTokens()) + .withStopSequences(fromOptions.getStopSequences()) + .withRepetitionPenalty(fromOptions.getRepetitionPenalty()) + .withRandomSeed(fromOptions.getRandomSeed()) + .withModel(fromOptions.getModel()) + .withAdditionalProperties(fromOptions.getAdditionalProperties()) + .build(); + } + } // @formatter:on \ No newline at end of file diff --git a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiModelCaller.java b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiModelCaller.java index 239f4ffc2..1dbf908ae 100644 --- a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiModelCaller.java +++ b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiModelCaller.java @@ -140,4 +140,9 @@ public class WatsonxAiModelCaller implements ChatCaller, StreamingChatCaller { return WatsonxAiRequest.builder(convertedPrompt).withParameters(parameters).build(); } + @Override + public ChatOptions getDefaultOptions() { + return WatsonxAiChatOptions.fromOptions(this.defaultOptions); + } + } \ No newline at end of file diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatClient.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatCaller.java similarity index 97% rename from models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatClient.java rename to models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatCaller.java index 39e86bdb4..d02f2a105 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatClient.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatCaller.java @@ -66,11 +66,11 @@ import java.util.concurrent.ConcurrentHashMap; * @see ZhiPuAiApi * @since 1.0.0 M1 */ -public class ZhiPuAiChatClient extends +public class ZhiPuAiChatCaller extends AbstractFunctionCallSupport> implements ChatCaller, StreamingChatCaller { - private static final Logger logger = LoggerFactory.getLogger(ZhiPuAiChatClient.class); + private static final Logger logger = LoggerFactory.getLogger(ZhiPuAiChatCaller.class); /** * The default options used for the chat completion requests. @@ -93,7 +93,7 @@ public class ZhiPuAiChatClient extends * ZhiPuAI Chat API. * @throws IllegalArgumentException if zhiPuAiApi is null */ - public ZhiPuAiChatClient(ZhiPuAiApi zhiPuAiApi) { + public ZhiPuAiChatCaller(ZhiPuAiApi zhiPuAiApi) { this(zhiPuAiApi, ZhiPuAiChatOptions.builder().withModel(ZhiPuAiApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); } @@ -104,7 +104,7 @@ public class ZhiPuAiChatClient extends * ZhiPuAI Chat API. * @param options The ZhiPuAiChatOptions to configure the chat client. */ - public ZhiPuAiChatClient(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options) { + public ZhiPuAiChatCaller(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options) { this(zhiPuAiApi, options, null, RetryUtils.DEFAULT_RETRY_TEMPLATE); } @@ -116,7 +116,7 @@ public class ZhiPuAiChatClient extends * @param functionCallbackContext The function callback context. * @param retryTemplate The retry template. */ - public ZhiPuAiChatClient(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options, + public ZhiPuAiChatCaller(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { super(functionCallbackContext); Assert.notNull(zhiPuAiApi, "ZhiPuAiApi must not be null"); @@ -381,4 +381,9 @@ public class ZhiPuAiChatClient extends && choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; } + @Override + public ChatOptions getDefaultOptions() { + return ZhiPuAiChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java index 13dc8d519..d7a8dc211 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java @@ -490,4 +490,24 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions { throw new UnsupportedOperationException("Unimplemented method 'setTopK'"); } + public static ZhiPuAiChatOptions fromOptions(ZhiPuAiChatOptions fromOptions) { + return ZhiPuAiChatOptions.builder() + .withModel(fromOptions.getModel()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withN(fromOptions.getN()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withResponseFormat(fromOptions.getResponseFormat()) + .withSeed(fromOptions.getSeed()) + .withStop(fromOptions.getStop()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTools(fromOptions.getTools()) + .withToolChoice(fromOptions.getToolChoice()) + .withUser(fromOptions.getUser()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ChatCompletionRequestTests.java b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ChatCompletionRequestTests.java index b958bc013..45b48d732 100644 --- a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ChatCompletionRequestTests.java +++ b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ChatCompletionRequestTests.java @@ -33,7 +33,7 @@ public class ChatCompletionRequestTests { @Test public void createRequestWithChatOptions() { - var client = new ZhiPuAiChatClient(new ZhiPuAiApi("TEST"), + var client = new ZhiPuAiChatCaller(new ZhiPuAiApi("TEST"), ZhiPuAiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); var request = client.createRequest(new Prompt("Test message content"), false); @@ -59,7 +59,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new ZhiPuAiChatClient(new ZhiPuAiApi("TEST"), + var client = new ZhiPuAiChatCaller(new ZhiPuAiApi("TEST"), ZhiPuAiChatOptions.builder().withModel("DEFAULT_MODEL").build()); var request = client.createRequest(new Prompt("Test message content", @@ -89,7 +89,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new ZhiPuAiChatClient(new ZhiPuAiApi("TEST"), + var client = new ZhiPuAiChatCaller(new ZhiPuAiApi("TEST"), ZhiPuAiChatOptions.builder() .withModel("DEFAULT_MODEL") .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) diff --git a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ZhiPuAiTestConfiguration.java b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ZhiPuAiTestConfiguration.java index 92f0bbb28..635204f20 100644 --- a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ZhiPuAiTestConfiguration.java +++ b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ZhiPuAiTestConfiguration.java @@ -48,8 +48,8 @@ public class ZhiPuAiTestConfiguration { } @Bean - public ZhiPuAiChatClient zhiPuAiChatClient(ZhiPuAiApi api) { - return new ZhiPuAiChatClient(api); + public ZhiPuAiChatCaller zhiPuAiChatClient(ZhiPuAiApi api) { + return new ZhiPuAiChatCaller(api); } @Bean diff --git a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java index 3326af098..2b3dbe2a8 100644 --- a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java +++ b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java @@ -26,7 +26,7 @@ import org.springframework.ai.image.ImageMessage; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.retry.RetryUtils; import org.springframework.ai.retry.TransientAiException; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatCaller; import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingClient; import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingOptions; @@ -93,7 +93,7 @@ public class ZhiPuAiRetryTests { private @Mock ZhiPuAiImageApi zhiPuAiImageApi; - private ZhiPuAiChatClient chatClient; + private ZhiPuAiChatCaller chatClient; private ZhiPuAiEmbeddingClient embeddingClient; @@ -105,7 +105,7 @@ public class ZhiPuAiRetryTests { retryListener = new TestRetryListener(); retryTemplate.registerListener(retryListener); - chatClient = new ZhiPuAiChatClient(zhiPuAiApi, ZhiPuAiChatOptions.builder().build(), null, retryTemplate); + chatClient = new ZhiPuAiChatCaller(zhiPuAiApi, ZhiPuAiChatOptions.builder().build(), null, retryTemplate); embeddingClient = new ZhiPuAiEmbeddingClient(zhiPuAiApi, MetadataMode.EMBED, ZhiPuAiEmbeddingOptions.builder().build(), retryTemplate); imageClient = new ZhiPuAiImageClient(zhiPuAiImageApi, ZhiPuAiImageOptions.builder().build(), retryTemplate); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatCaller.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatCaller.java index 8498801ac..12342027f 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatCaller.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatCaller.java @@ -15,6 +15,7 @@ */ package org.springframework.ai.chat; +import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import java.util.Arrays; @@ -23,7 +24,7 @@ import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.model.ModelClient; -@FunctionalInterface +// @FunctionalInterface public interface ChatCaller extends ModelClient { default String call(String message) { @@ -41,4 +42,6 @@ public interface ChatCaller extends ModelClient { @Override ChatResponse call(Prompt prompt); + ChatOptions getDefaultOptions(); + } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java index 9a5979266..82fa65d7c 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java @@ -29,7 +29,6 @@ import java.util.function.Consumer; import reactor.core.publisher.Flux; -import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.Media; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.SystemMessage; @@ -51,7 +50,7 @@ import org.springframework.util.StringUtils; // todo support plugging in a outputConverter at runtime // todo figure out stream and list methods -/* +/** * @author Mark Pollack * @author Christian Tzolov * @author Josh Long @@ -206,7 +205,7 @@ public interface ChatClient { ChatOptions chatOptions) { this.caller = caller; - this.chatOptions = chatOptions; + this.chatOptions = chatOptions != null ? chatOptions : caller.getDefaultOptions(); this.userText = userText; this.systemText = systemText; @@ -229,10 +228,10 @@ public interface ChatClient { public ChatClientRequest function(String name, String description, java.util.function.Function function) { var fcw = FunctionCallbackWrapper.builder(function) - .withDescription(description) - .withName(name) - .withResponseConverter(Object::toString) - .build(); + .withDescription(description) + .withName(name) + .withResponseConverter(Object::toString) + .build(); this.functionCallbacks.add(fcw); return this; } @@ -250,27 +249,32 @@ public interface ChatClient { public ChatClientRequest system(Consumer consumer) { var ss = new SystemSpec(); consumer.accept(ss); - this.systemText = ss.text(); + this.systemText = StringUtils.hasText(ss.text()) ? ss.text() : this.systemText; this.systemParams.putAll(ss.params()); return this; } + public ChatClientRequest user(String text) { + this.userText = text; + return this; + } + public ChatClientRequest user(Consumer consumer) { var us = new UserSpec(); consumer.accept(us); - this.userText = us.text(); + this.userText = StringUtils.hasText(us.text()) ? us.text() : this.userText; this.userParams.putAll(us.params()); this.media.addAll(us.media()); return this; } - public static class ChatResponseSpec { + public static class CollectResponseSpec { private final ChatClientRequest request; private final ChatCaller modelCaller; - public ChatResponseSpec(ChatCaller modelCaller, ChatClientRequest request) { + public CollectResponseSpec(ChatCaller modelCaller, ChatClientRequest request) { this.modelCaller = modelCaller; this.request = request; } @@ -353,8 +357,10 @@ public interface ChatClient { } public List contents() { - return doGetChatResponse(this.request.userText).getResults().stream() - .map(r -> r.getOutput().getContent()).toList(); + return doGetChatResponse(this.request.userText).getResults() + .stream() + .map(r -> r.getOutput().getContent()) + .toList(); } @SuppressWarnings("unused") @@ -369,24 +375,27 @@ public interface ChatClient { } - public static class ChatStreamResponseSpec { + public static class StreamResponseSpec { private final ChatClientRequest request; private final StreamingChatCaller modelCaller; - public ChatStreamResponseSpec(StreamingChatCaller modelCaller, ChatClientRequest request) { + public StreamResponseSpec(StreamingChatCaller modelCaller, ChatClientRequest request) { this.modelCaller = modelCaller; this.request = request; } // public Flux single(ParameterizedTypeReference t) { - // return doSingleWithBeanOutputConverter(new BeanOutputConverter(new ParameterizedTypeReference<>() { + // return doSingleWithBeanOutputConverter(new BeanOutputConverter(new + // ParameterizedTypeReference<>() { // })); // } - // private Flux doSingleWithBeanOutputConverter(BeanOutputConverter boc) { - // var processedUserText = this.request.userText + System.lineSeparator() + System.lineSeparator() + // private Flux doSingleWithBeanOutputConverter(BeanOutputConverter + // boc) { + // var processedUserText = this.request.userText + System.lineSeparator() + + // System.lineSeparator() // + "{format}"; // var chatResponse = doGetChatResponse(processedUserText, boc.getFormat()); // var stringResponse = chatResponse.getResult().getOutput().getContent(); @@ -399,15 +408,13 @@ public interface ChatClient { // return doSingleWithBeanOutputConverter(boc); // } - private Flux doGetChatResponse(String processedUserText) { - return this.doGetChatResponse(processedUserText, ""); - } + // private Flux doGetFluxChatResponse(String processedUserText) + // { + // return this.doGetFluxChatResponse(processedUserText, ""); + // } - private Flux doGetChatResponse(String processedUserText, String formatParam) { + private Flux doGetFluxChatResponse(String processedUserText) { Map userParams = new HashMap<>(this.request.userParams); - if (StringUtils.hasText(formatParam)) { - userParams.put("format", formatParam); - } var messages = new ArrayList(); var textsAreValid = (StringUtils.hasText(processedUserText) @@ -450,21 +457,22 @@ public interface ChatClient { } public Flux chatResponse() { - return doGetChatResponse(this.request.userText); + return doGetFluxChatResponse(this.request.userText); } public Flux content() { - return doGetChatResponse(this.request.userText) - .map(ChatResponse::getResult) - .map(Generation::getOutput) - .map(AssistantMessage::getContent); - // .map(r -> r.getResult().getOutput().getContent()) - // .filter(v -> StringUtils.hasText(v)); - } - - public Flux> contents() { - return doGetChatResponse(this.request.userText).map(r -> r.getResults().stream() - .map(rr -> rr.getOutput().getContent()).toList()); + return doGetFluxChatResponse(this.request.userText) + // .map(ChatResponse::getResult) + // .map(Generation::getOutput) + // .map(AssistantMessage::getContent); + .map(r -> { + if (r.getResult() == null || r.getResult().getOutput() == null + || r.getResult().getOutput().getContent() == null) { + return ""; + } + return r.getResult().getOutput().getContent(); + }) + .filter(v -> StringUtils.hasText(v)); } // @SuppressWarnings("unused") @@ -479,12 +487,12 @@ public interface ChatClient { } - public ChatResponseSpec chat() { - return new ChatResponseSpec(this.caller, this); + public CollectResponseSpec collect() { + return new CollectResponseSpec(this.caller, this); } - public ChatStreamResponseSpec stream() { - return new ChatStreamResponseSpec((StreamingChatCaller) this.caller, this); + public StreamResponseSpec stream() { + return new StreamResponseSpec((StreamingChatCaller) this.caller, this); } } @@ -505,7 +513,7 @@ public interface ChatClient { return new DefaultChatClient(this.modelCaller, this.defaultRequest); } - public ChatClientBuilder defaultChatOptions(ChatOptions chatOptions) { + public ChatClientBuilder defaultRuntimeOptions(ChatOptions chatOptions) { this.defaultRequest.chatOptions(chatOptions); return this; } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/Main.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/Main.java index 120a40eed..ac2e143cc 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/Main.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/Main.java @@ -17,7 +17,6 @@ package org.springframework.ai.chat; import java.net.MalformedURLException; -import java.net.URI; import java.net.URL; import org.junit.jupiter.api.Test; @@ -25,8 +24,6 @@ import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; -import org.springframework.core.io.DefaultResourceLoader; -import org.springframework.util.MimeType; import org.springframework.util.MimeTypeUtils; /** @@ -51,7 +48,7 @@ public class Main { String response = client.call() .user(u -> u.text("User text {music}").param("music", "Rock").media(MimeTypeUtils.IMAGE_PNG, url)) - .chat() + .collect() .single(String.class); } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfiguration.java index f3c805db2..d257fb7c1 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfiguration.java @@ -16,7 +16,7 @@ package org.springframework.ai.autoconfigure.minimax; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatCaller; import org.springframework.ai.minimax.MiniMaxEmbeddingClient; import org.springframework.ai.minimax.api.MiniMaxApi; import org.springframework.ai.model.function.FunctionCallback; @@ -51,7 +51,7 @@ public class MiniMaxAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = MiniMaxChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public MiniMaxChatClient miniMaxChatClient(MiniMaxConnectionProperties commonProperties, + public MiniMaxChatCaller miniMaxChatClient(MiniMaxConnectionProperties commonProperties, MiniMaxChatProperties chatProperties, RestClient.Builder restClientBuilder, List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -63,7 +63,7 @@ public class MiniMaxAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - return new MiniMaxChatClient(miniMaxApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); + return new MiniMaxChatCaller(miniMaxApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); } @Bean diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfiguration.java index e2c8b1126..9b1788d92 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfiguration.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.zhipuai; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatCaller; import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingClient; import org.springframework.ai.zhipuai.ZhiPuAiImageClient; import org.springframework.ai.zhipuai.api.ZhiPuAiApi; @@ -53,7 +53,7 @@ public class ZhiPuAiAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = ZhiPuAiChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public ZhiPuAiChatClient zhiPuAiChatClient(ZhiPuAiConnectionProperties commonProperties, + public ZhiPuAiChatCaller zhiPuAiChatClient(ZhiPuAiConnectionProperties commonProperties, ZhiPuAiChatProperties chatProperties, RestClient.Builder restClientBuilder, List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -65,7 +65,7 @@ public class ZhiPuAiAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - return new ZhiPuAiChatClient(zhiPuAiApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); + return new ZhiPuAiChatCaller(zhiPuAiApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); } @Bean diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackInPromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackInPromptIT.java index 7026812a6..9eb7c110b 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackInPromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackInPromptIT.java @@ -25,7 +25,7 @@ import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatCaller; import org.springframework.ai.minimax.MiniMaxChatOptions; import org.springframework.ai.model.function.FunctionCallbackWrapper; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -55,7 +55,7 @@ public class FunctionCallbackInPromptIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller chatClient = context.getBean(MiniMaxChatCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -80,7 +80,7 @@ public class FunctionCallbackInPromptIT { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller chatClient = context.getBean(MiniMaxChatCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWithPlainFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWithPlainFunctionBeanIT.java index 1c1492e65..6d4cacd65 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWithPlainFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWithPlainFunctionBeanIT.java @@ -25,7 +25,7 @@ import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatCaller; import org.springframework.ai.minimax.MiniMaxChatOptions; import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; @@ -61,7 +61,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller chatClient = context.getBean(MiniMaxChatCaller.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -88,7 +88,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallWithPortableFunctionCallingOptions() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller chatClient = context.getBean(MiniMaxChatCaller.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -107,7 +107,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller chatClient = context.getBean(MiniMaxChatCaller.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWrapperIT.java index 75530376f..11ceae1b1 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWrapperIT.java @@ -25,7 +25,7 @@ import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatCaller; import org.springframework.ai.minimax.MiniMaxChatOptions; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackWrapper; @@ -59,7 +59,7 @@ public class FunctionCallbackWrapperIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller chatClient = context.getBean(MiniMaxChatCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -77,7 +77,7 @@ public class FunctionCallbackWrapperIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller chatClient = context.getBean(MiniMaxChatCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfigurationIT.java index d400b2c47..b5bf46a99 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfigurationIT.java @@ -24,7 +24,7 @@ import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatCaller; import org.springframework.ai.minimax.MiniMaxEmbeddingClient; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -52,7 +52,7 @@ public class MiniMaxAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - MiniMaxChatClient client = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller client = context.getBean(MiniMaxChatCaller.class); String response = client.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); @@ -62,7 +62,7 @@ public class MiniMaxAutoConfigurationIT { @Test void generateStreaming() { contextRunner.run(context -> { - MiniMaxChatClient client = context.getBean(MiniMaxChatClient.class); + MiniMaxChatCaller client = context.getBean(MiniMaxChatCaller.class); Flux responseFlux = client.stream(new Prompt(new UserMessage("Hello"))); String response = responseFlux.collectList().block().stream().map(chatResponse -> { return chatResponse.getResults().get(0).getOutput().getContent(); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxPropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxPropertiesTests.java index 31e103508..a047fac92 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxPropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxPropertiesTests.java @@ -19,7 +19,7 @@ import org.junit.jupiter.api.Test; import org.skyscreamer.jsonassert.JSONAssert; import org.skyscreamer.jsonassert.JSONCompareMode; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatCaller; import org.springframework.ai.minimax.MiniMaxEmbeddingClient; import org.springframework.ai.minimax.api.MiniMaxApi; import org.springframework.ai.model.ModelOptionsUtils; @@ -302,7 +302,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(MiniMaxChatCaller.class)).isEmpty(); }); new ApplicationContextRunner() @@ -311,7 +311,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(MiniMaxChatCaller.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -321,7 +321,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(MiniMaxChatCaller.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java new file mode 100644 index 000000000..e78c99e11 --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java @@ -0,0 +1,115 @@ +/* + * Copyright 2023 - 2024 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.autoconfigure.openai.tool; + +import java.util.stream.Collectors; + +import org.junit.jupiter.api.Test; +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.autoconfigure.retry.SpringAiRetryAutoConfiguration; +import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.model.function.FunctionCallback; +import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.openai.OpenAiModelCaller; +import org.springframework.boot.autoconfigure.AutoConfigurations; +import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; +import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; + +import static org.assertj.core.api.Assertions.assertThat; + +@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".*") +public class FunctionCallbackWrapper2IT { + + private final Logger logger = LoggerFactory.getLogger(FunctionCallbackWrapperIT.class); + + private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() + .withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY")) + .withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class, + RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) + .withUserConfiguration(Config.class); + + @Test + void functionCallTest() { + contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { + + OpenAiModelCaller caller = context.getBean(OpenAiModelCaller.class); + + ChatClient chatClient = ChatClient.builder(caller) + .defaultFunctions("WeatherInfo") + .defaultUser(u -> u.text("What's the weather like in {cities}?")) + .build(); + + String content = chatClient.call() + .user(u -> u.param("cities", "San Francisco, Tokyo, Paris")) + .collect() + .content(); + + logger.info("Response: {}", content); + + assertThat(content).containsAnyOf("30.0", "30"); + assertThat(content).containsAnyOf("15.0", "15"); + assertThat(content).containsAnyOf("10", "10"); + }); + } + + @Test + void streamFunctionCallTest() { + contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { + + OpenAiModelCaller caller = context.getBean(OpenAiModelCaller.class); + + String content = ChatClient.builder(caller) + .build() + .call() + .functions("WeatherInfo") + .user(u -> u.text("What's the weather like in San Francisco, Tokyo, and Paris?")) + .stream() + .content() + .collectList() + .block() + .stream() + .collect(Collectors.joining()); + + logger.info("Response: {}", content); + + assertThat(content).containsAnyOf("30.0", "30"); + assertThat(content).containsAnyOf("10.0", "10"); + assertThat(content).containsAnyOf("15.0", "15"); + }); + } + + @Configuration + static class Config { + + @Bean + public FunctionCallback weatherFunctionInfo() { + + return FunctionCallbackWrapper.builder(new MockWeatherService()) + .withName("WeatherInfo") + .withDescription("Get the weather in location") + .withResponseConverter((response) -> "" + response.temp() + response.unit()) + .build(); + } + + } + +} \ No newline at end of file diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapperIT.java index 00e414061..2e6f6edde 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapperIT.java @@ -58,11 +58,11 @@ public class FunctionCallbackWrapperIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiModelCaller chatClient = context.getBean(OpenAiModelCaller.class); + OpenAiModelCaller caller = context.getBean(OpenAiModelCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - ChatResponse response = chatClient.call( + ChatResponse response = caller.call( new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("WeatherInfo").build())); logger.info("Response: {}", response); @@ -76,11 +76,11 @@ public class FunctionCallbackWrapperIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiModelCaller chatClient = context.getBean(OpenAiModelCaller.class); + OpenAiModelCaller caller = context.getBean(OpenAiModelCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - Flux response = chatClient.stream( + Flux response = caller.stream( new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("WeatherInfo").build())); String content = response.collectList() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfigurationIT.java index 073c81ffc..6bc0873d2 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfigurationIT.java @@ -26,7 +26,7 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.image.ImageResponse; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatCaller; import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingClient; import org.springframework.ai.zhipuai.ZhiPuAiImageClient; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -55,7 +55,7 @@ public class ZhiPuAiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - ZhiPuAiChatClient client = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller client = context.getBean(ZhiPuAiChatCaller.class); String response = client.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); @@ -65,7 +65,7 @@ public class ZhiPuAiAutoConfigurationIT { @Test void generateStreaming() { contextRunner.run(context -> { - ZhiPuAiChatClient client = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller client = context.getBean(ZhiPuAiChatCaller.class); Flux responseFlux = client.stream(new Prompt(new UserMessage("Hello"))); String response = responseFlux.collectList().block().stream().map(chatResponse -> { return chatResponse.getResults().get(0).getOutput().getContent(); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiPropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiPropertiesTests.java index 121473dfd..2b4715a67 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiPropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiPropertiesTests.java @@ -20,7 +20,7 @@ import org.skyscreamer.jsonassert.JSONAssert; import org.skyscreamer.jsonassert.JSONCompareMode; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.model.ModelOptionsUtils; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatCaller; import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingClient; import org.springframework.ai.zhipuai.ZhiPuAiImageClient; import org.springframework.ai.zhipuai.api.ZhiPuAiApi; @@ -378,7 +378,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiChatCaller.class)).isEmpty(); }); new ApplicationContextRunner() @@ -387,7 +387,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiChatCaller.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -397,7 +397,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiChatCaller.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java index 29a71041f..18264f9b1 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java @@ -27,7 +27,7 @@ import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallbackWrapper; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatCaller; import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -56,7 +56,7 @@ public class FunctionCallbackInPromptIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller chatClient = context.getBean(ZhiPuAiChatCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -81,7 +81,7 @@ public class FunctionCallbackInPromptIT { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller chatClient = context.getBean(ZhiPuAiChatCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWithPlainFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWithPlainFunctionBeanIT.java index 2d7c121d1..f118c5326 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWithPlainFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWithPlainFunctionBeanIT.java @@ -28,7 +28,7 @@ import org.springframework.ai.chat.messages.UserMessage; 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.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatCaller; import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -62,7 +62,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller chatClient = context.getBean(ZhiPuAiChatCaller.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -89,7 +89,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallWithPortableFunctionCallingOptions() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller chatClient = context.getBean(ZhiPuAiChatCaller.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -108,7 +108,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller chatClient = context.getBean(ZhiPuAiChatCaller.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWrapperIT.java index fefd81326..8e853f8b3 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWrapperIT.java @@ -28,7 +28,7 @@ import org.springframework.ai.chat.messages.UserMessage; 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.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatCaller; import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -60,7 +60,7 @@ public class FunctionCallbackWrapperIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller chatClient = context.getBean(ZhiPuAiChatCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -78,7 +78,7 @@ public class FunctionCallbackWrapperIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatCaller chatClient = context.getBean(ZhiPuAiChatCaller.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?");