From 7e2cefc3aa532a277ec60da3836fc2ff1a573aeb Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Thu, 19 Dec 2024 11:23:01 -0500 Subject: [PATCH] Share options instance between DefaultFunctionCallingOptionsBuilder and parent class --- .../prompt/DefaultChatOptionsBuilder.java | 2 +- .../DefaultFunctionCallingOptionsBuilder.java | 22 +-- ...ultFunctionCallingOptionsBuilderTests.java | 125 +++++++++++++++++- 3 files changed, 139 insertions(+), 10 deletions(-) diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/DefaultChatOptionsBuilder.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/DefaultChatOptionsBuilder.java index 1d84a7043..6956c3e8c 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/DefaultChatOptionsBuilder.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/DefaultChatOptionsBuilder.java @@ -23,7 +23,7 @@ import java.util.List; */ public class DefaultChatOptionsBuilder> implements ChatOptions.Builder { - private final DefaultChatOptions options = new DefaultChatOptions(); + protected DefaultChatOptions options = new DefaultChatOptions(); protected T self() { return (T) this; diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/DefaultFunctionCallingOptionsBuilder.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/DefaultFunctionCallingOptionsBuilder.java index 3d6d12415..590e91b99 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/DefaultFunctionCallingOptionsBuilder.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/DefaultFunctionCallingOptionsBuilder.java @@ -36,22 +36,28 @@ public class DefaultFunctionCallingOptionsBuilder extends DefaultChatOptionsBuilder implements FunctionCallingOptions.Builder { - private final DefaultFunctionCallingOptions functionCallingOptions = new DefaultFunctionCallingOptions(); + private final DefaultFunctionCallingOptions functionCallingOptions; + + public DefaultFunctionCallingOptionsBuilder() { + this.functionCallingOptions = new DefaultFunctionCallingOptions(); + // Set the options in the parent class to be the same instance + super.options = this.functionCallingOptions; + } public DefaultFunctionCallingOptionsBuilder functionCallbacks(List functionCallbacks) { this.functionCallingOptions.setFunctionCallbacks(functionCallbacks); - return this; + return self(); } public DefaultFunctionCallingOptionsBuilder functionCallbacks(FunctionCallback... functionCallbacks) { Assert.notNull(functionCallbacks, "FunctionCallbacks must not be null"); this.functionCallingOptions.setFunctionCallbacks(List.of(functionCallbacks)); - return this; + return self(); } public DefaultFunctionCallingOptionsBuilder functions(Set functions) { this.functionCallingOptions.setFunctions(functions); - return this; + return self(); } public DefaultFunctionCallingOptionsBuilder function(String function) { @@ -59,12 +65,12 @@ public class DefaultFunctionCallingOptionsBuilder var set = new HashSet<>(this.functionCallingOptions.getFunctions()); set.add(function); this.functionCallingOptions.setFunctions(set); - return this; + return self(); } public DefaultFunctionCallingOptionsBuilder proxyToolCalls(Boolean proxyToolCalls) { this.functionCallingOptions.setProxyToolCalls(proxyToolCalls); - return this; + return self(); } public DefaultFunctionCallingOptionsBuilder toolContext(Map context) { @@ -72,7 +78,7 @@ public class DefaultFunctionCallingOptionsBuilder Map newContext = new HashMap<>(this.functionCallingOptions.getToolContext()); newContext.putAll(context); this.functionCallingOptions.setToolContext(newContext); - return this; + return self(); } public DefaultFunctionCallingOptionsBuilder toolContext(String key, Object value) { @@ -81,7 +87,7 @@ public class DefaultFunctionCallingOptionsBuilder Map newContext = new HashMap<>(this.functionCallingOptions.getToolContext()); newContext.put(key, value); this.functionCallingOptions.setToolContext(newContext); - return this; + return self(); } public FunctionCallingOptions build() { diff --git a/spring-ai-core/src/test/java/org/springframework/ai/model/function/DefaultFunctionCallingOptionsBuilderTests.java b/spring-ai-core/src/test/java/org/springframework/ai/model/function/DefaultFunctionCallingOptionsBuilderTests.java index be603f032..7cc828cd7 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/model/function/DefaultFunctionCallingOptionsBuilderTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/model/function/DefaultFunctionCallingOptionsBuilderTests.java @@ -25,6 +25,8 @@ import java.util.List; import java.util.Map; import java.util.Set; +import org.springframework.ai.chat.prompt.ChatOptions; + import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; @@ -41,6 +43,109 @@ class DefaultFunctionCallingOptionsBuilderTests { builder = new DefaultFunctionCallingOptionsBuilder(); } + // Tests for inherited ChatOptions properties + + @Test + void shouldBuildWithModel() { + // When + ChatOptions options = builder.model("gpt-4").build(); + + // Then + assertThat(options.getModel()).isEqualTo("gpt-4"); + } + + @Test + void shouldBuildWithFrequencyPenalty() { + // When + ChatOptions options = builder.frequencyPenalty(0.5).build(); + + // Then + assertThat(options.getFrequencyPenalty()).isEqualTo(0.5); + } + + @Test + void shouldBuildWithMaxTokens() { + // When + ChatOptions options = builder.maxTokens(100).build(); + + // Then + assertThat(options.getMaxTokens()).isEqualTo(100); + } + + @Test + void shouldBuildWithPresencePenalty() { + // When + ChatOptions options = builder.presencePenalty(0.7).build(); + + // Then + assertThat(options.getPresencePenalty()).isEqualTo(0.7); + } + + @Test + void shouldBuildWithStopSequences() { + // Given + List stopSequences = List.of("stop1", "stop2"); + + // When + ChatOptions options = builder.stopSequences(stopSequences).build(); + + // Then + assertThat(options.getStopSequences()).hasSize(2).containsExactlyElementsOf(stopSequences); + } + + @Test + void shouldBuildWithTemperature() { + // When + ChatOptions options = builder.temperature(0.8).build(); + + // Then + assertThat(options.getTemperature()).isEqualTo(0.8); + } + + @Test + void shouldBuildWithTopK() { + // When + ChatOptions options = builder.topK(5).build(); + + // Then + assertThat(options.getTopK()).isEqualTo(5); + } + + @Test + void shouldBuildWithTopP() { + // When + ChatOptions options = builder.topP(0.9).build(); + + // Then + assertThat(options.getTopP()).isEqualTo(0.9); + } + + @Test + void shouldBuildWithAllInheritedOptions() { + // When + ChatOptions options = builder.model("gpt-4") + .frequencyPenalty(0.5) + .maxTokens(100) + .presencePenalty(0.7) + .stopSequences(List.of("stop1", "stop2")) + .temperature(0.8) + .topK(5) + .topP(0.9) + .build(); + + // Then + assertThat(options.getModel()).isEqualTo("gpt-4"); + assertThat(options.getFrequencyPenalty()).isEqualTo(0.5); + assertThat(options.getMaxTokens()).isEqualTo(100); + assertThat(options.getPresencePenalty()).isEqualTo(0.7); + assertThat(options.getStopSequences()).containsExactly("stop1", "stop2"); + assertThat(options.getTemperature()).isEqualTo(0.8); + assertThat(options.getTopK()).isEqualTo(5); + assertThat(options.getTopP()).isEqualTo(0.9); + } + + // Original FunctionCallingOptions tests + @Test void shouldBuildWithFunctionCallbacksList() { // Given @@ -195,7 +300,15 @@ class DefaultFunctionCallingOptionsBuilderTests { Map context = Map.of("key1", "value1"); // When - FunctionCallingOptions options = builder.functionCallbacks(callback) + FunctionCallingOptions options = builder.model("gpt-4") + .frequencyPenalty(0.5) + .maxTokens(100) + .presencePenalty(0.7) + .stopSequences(List.of("stop1", "stop2")) + .temperature(0.8) + .topK(5) + .topP(0.9) + .functionCallbacks(callback) .functions(functions) .proxyToolCalls(true) .toolContext(context) @@ -206,6 +319,16 @@ class DefaultFunctionCallingOptionsBuilderTests { assertThat(options.getFunctions()).hasSize(1).containsExactlyElementsOf(functions); assertThat(options.getProxyToolCalls()).isTrue(); assertThat(options.getToolContext()).hasSize(1).containsAllEntriesOf(context); + + ChatOptions chatOptions = options; + assertThat(chatOptions.getModel()).isEqualTo("gpt-4"); + assertThat(chatOptions.getFrequencyPenalty()).isEqualTo(0.5); + assertThat(chatOptions.getMaxTokens()).isEqualTo(100); + assertThat(chatOptions.getPresencePenalty()).isEqualTo(0.7); + assertThat(chatOptions.getStopSequences()).containsExactly("stop1", "stop2"); + assertThat(chatOptions.getTemperature()).isEqualTo(0.8); + assertThat(chatOptions.getTopK()).isEqualTo(5); + assertThat(chatOptions.getTopP()).isEqualTo(0.9); } }