From 3938cc83ad823caa0fd6627f5e8a6a710a22c5db Mon Sep 17 00:00:00 2001 From: youngpar Date: Wed, 6 Mar 2024 21:55:15 +0900 Subject: [PATCH] Remove setters from options interface * Add test code --- .../azure/openai/AzureOpenAiChatOptions.java | 5 - .../anthropic/AnthropicChatOptions.java | 3 - .../cohere/BedrockCohereChatOptions.java | 3 - .../llama2/BedrockLlama2ChatOptions.java | 1 - .../titan/BedrockTitanChatOptions.java | 1 - .../ai/mistralai/MistralAiChatOptions.java | 3 - .../ai/openai/OpenAiChatOptions.java | 4 - .../gemini/VertexAiGeminiChatOptions.java | 3 - .../palm2/VertexAiPaLm2ChatOptions.java | 3 - .../ai/chat/prompt/ChatOptions.java | 6 - .../ai/chat/prompt/ChatOptionsBuilder.java | 5 +- .../function/FunctionCallingOptions.java | 2 +- .../FunctionCallingOptionsBuilder.java | 7 +- .../ai/chat/ChatBuilderTests.java | 111 ++++++++++++++++++ 14 files changed, 114 insertions(+), 43 deletions(-) create mode 100644 spring-ai-core/src/test/java/org/springframework/ai/chat/ChatBuilderTests.java 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 c10636f6e..2d7ee90ad 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 @@ -312,7 +312,6 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio return this.temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -322,7 +321,6 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio return this.topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } @@ -333,7 +331,6 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio throw new UnsupportedOperationException("Unimplemented method 'getTopK'"); } - @Override @JsonIgnore public void setTopK(Integer topK) { throw new UnsupportedOperationException("Unimplemented method 'setTopK'"); @@ -344,7 +341,6 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio return this.functionCallbacks; } - @Override public void setFunctionCallbacks(List functionCallbacks) { this.functionCallbacks = functionCallbacks; } @@ -354,7 +350,6 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio return this.functions; } - @Override public void setFunctions(Set functions) { this.functions = functions; } 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 75e209b6a..70262ab42 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 @@ -119,7 +119,6 @@ public class AnthropicChatOptions implements ChatOptions { return this.temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -137,7 +136,6 @@ public class AnthropicChatOptions implements ChatOptions { return this.topK; } - @Override public void setTopK(Integer topK) { this.topK = topK; } @@ -147,7 +145,6 @@ public class AnthropicChatOptions implements ChatOptions { return this.topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } 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 d1db43ac7..aa3e727bb 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 @@ -144,7 +144,6 @@ public class BedrockCohereChatOptions implements ChatOptions { return this.temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -154,7 +153,6 @@ public class BedrockCohereChatOptions implements ChatOptions { return this.topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } @@ -164,7 +162,6 @@ public class BedrockCohereChatOptions implements ChatOptions { return this.topK; } - @Override public void setTopK(Integer topK) { this.topK = topK; } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama2/BedrockLlama2ChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama2/BedrockLlama2ChatOptions.java index 5c2bed1ef..b5c14a866 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama2/BedrockLlama2ChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama2/BedrockLlama2ChatOptions.java @@ -105,7 +105,6 @@ public class BedrockLlama2ChatOptions implements ChatOptions { throw new UnsupportedOperationException("Unsupported option: 'TopK'"); } - @Override @JsonIgnore public void setTopK(Integer topK) { throw new UnsupportedOperationException("Unsupported option: 'TopK'"); 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 86d93fd76..e461692ab 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 @@ -125,7 +125,6 @@ public class BedrockTitanChatOptions implements ChatOptions { throw new UnsupportedOperationException("Bedrock Titian Chat does not support the 'TopK' option."); } - @Override public void setTopK(Integer topK) { throw new UnsupportedOperationException("Bedrock Titian Chat does not support the 'TopK' option.'"); } 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 c47d95946..6527d217e 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 @@ -264,7 +264,6 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions return this.temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -274,7 +273,6 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions return this.topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } @@ -285,7 +283,6 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions throw new UnsupportedOperationException("Unsupported option: 'TopK'"); } - @Override @JsonIgnore public void setTopK(Integer topK) { throw new UnsupportedOperationException("Unsupported option: 'TopK'"); 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 3ebb1c9db..35f2d954a 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 @@ -333,7 +333,6 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions { return this.temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -343,7 +342,6 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions { return this.topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } @@ -387,7 +385,6 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions { return functions; } - @Override public void setFunctions(Set functionNames) { this.functions = functionNames; } @@ -515,7 +512,6 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions { throw new UnsupportedOperationException("Unimplemented method 'getTopK'"); } - @Override @JsonIgnore public void setTopK(Integer topK) { throw new UnsupportedOperationException("Unimplemented method 'setTopK'"); 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 2c78be4ce..ba0365548 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 @@ -190,7 +190,6 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp return this.temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -200,7 +199,6 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp return this.topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } @@ -215,7 +213,6 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp this.topK = topK; } - @Override @JsonIgnore public void setTopK(Integer topK) { this.topK = (topK != null) ? topK.floatValue() : null; 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 2683a0152..3ce296574 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 @@ -98,7 +98,6 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions { return this.temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -116,7 +115,6 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions { return this.topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } @@ -126,7 +124,6 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions { return this.topK; } - @Override public void setTopK(Integer topK) { this.topK = topK; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/ChatOptions.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/ChatOptions.java index 5d91bbd02..4abb8e1e7 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/ChatOptions.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/ChatOptions.java @@ -25,14 +25,8 @@ public interface ChatOptions extends ModelOptions { Float getTemperature(); - void setTemperature(Float temperature); - Float getTopP(); - void setTopP(Float topP); - Integer getTopK(); - void setTopK(Integer topK); - } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/ChatOptionsBuilder.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/ChatOptionsBuilder.java index f702f6350..1d49ab4b5 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/ChatOptionsBuilder.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/ChatOptionsBuilder.java @@ -31,7 +31,6 @@ public class ChatOptionsBuilder { return temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -41,7 +40,6 @@ public class ChatOptionsBuilder { return topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } @@ -51,7 +49,6 @@ public class ChatOptionsBuilder { return topK; } - @Override public void setTopK(Integer topK) { this.topK = topK; } @@ -86,4 +83,4 @@ public class ChatOptionsBuilder { return options; } -} +} \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java index c66a4f5b1..dc2329830 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java @@ -63,4 +63,4 @@ public interface FunctionCallingOptions { return new FunctionCallingOptionsBuilder(); } -} +} \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptionsBuilder.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptionsBuilder.java index 948fba58f..a77e5e0d1 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptionsBuilder.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptionsBuilder.java @@ -98,7 +98,6 @@ public class FunctionCallingOptionsBuilder { return this.functionCallbacks; } - @Override public void setFunctionCallbacks(List functionCallbacks) { Assert.notNull(functionCallbacks, "FunctionCallbacks must not be null"); this.functionCallbacks = functionCallbacks; @@ -109,7 +108,6 @@ public class FunctionCallingOptionsBuilder { return this.functions; } - @Override public void setFunctions(Set functions) { Assert.notNull(functions, "Functions must not be null"); this.functions = functions; @@ -120,7 +118,6 @@ public class FunctionCallingOptionsBuilder { return this.temperature; } - @Override public void setTemperature(Float temperature) { this.temperature = temperature; } @@ -130,7 +127,6 @@ public class FunctionCallingOptionsBuilder { return this.topP; } - @Override public void setTopP(Float topP) { this.topP = topP; } @@ -140,11 +136,10 @@ public class FunctionCallingOptionsBuilder { return this.topK; } - @Override public void setTopK(Integer topK) { this.topK = topK; } } -} +} \ No newline at end of file diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatBuilderTests.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatBuilderTests.java new file mode 100644 index 000000000..a123c4e9d --- /dev/null +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatBuilderTests.java @@ -0,0 +1,111 @@ +/* + * Copyright 2023 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.chat; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.ArrayList; +import java.util.HashSet; +import java.util.List; +import java.util.Set; +import org.junit.jupiter.api.Test; + +import org.springframework.ai.chat.prompt.ChatOptions; +import org.springframework.ai.chat.prompt.ChatOptionsBuilder; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.function.FunctionCallback; +import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.model.function.FunctionCallingOptions; +import org.springframework.ai.model.function.FunctionCallingOptionsBuilder; + +/** + * Unit Tests for {@link Prompt}. + * + * @author youngmon + * @since 0.8.1 + */ +public class ChatBuilderTests { + + @Test + void createNewChatOptionsTest() { + Float temperature = 1.1f; + Float topP = 2.2f; + Integer topK = 111; + + ChatOptions options = ChatOptionsBuilder.builder() + .withTemperature(temperature) + .withTopK(topK) + .withTopP(topP) + .build(); + + assertThat(options.getTemperature()).isEqualTo(temperature); + assertThat(options.getTopP()).isEqualTo(topP); + assertThat(options.getTopK()).isEqualTo(topK); + } + + @Test + void duplicateChatOptionsTest() { + Float initTemperature = 1.1f; + Float initTopP = 2.2f; + Integer initTopK = 111; + + ChatOptions options = ChatOptionsBuilder.builder() + .withTemperature(initTemperature) + .withTopP(initTopP) + .withTopK(initTopK) + .build(); + + } + + @Test + void createFunctionCallingOptionTest() { + Float temperature = 1.1f; + Float topP = 2.2f; + Integer topK = 111; + List functionCallbacks = new ArrayList<>(); + Set functions = new HashSet<>(); + + String func = "func"; + FunctionCallback cb = FunctionCallbackWrapper.builder(i -> i) + .withName("cb") + .withDescription("cb") + .build(); + + functions.add(func); + functionCallbacks.add(cb); + + FunctionCallingOptions options = FunctionCallingOptions.builder() + .withFunctionCallbacks(functionCallbacks) + .withFunctions(functions) + .withTopK(topK) + .withTopP(topP) + .withTemperature(temperature) + .build(); + + // Callback Functions + assertThat(options.getFunctionCallbacks()).isNotNull(); + assertThat(options.getFunctionCallbacks().size()).isEqualTo(1); + assertThat(options.getFunctionCallbacks().contains(cb)); + + // Functions + assertThat(options.getFunctions()).isNotNull(); + assertThat(options.getFunctions().size()).isEqualTo(1); + assertThat(options.getFunctions().contains(func)); + + } + +}