Share options instance between DefaultFunctionCallingOptionsBuilder and parent class

This commit is contained in:
Mark Pollack
2024-12-19 11:23:01 -05:00
parent adefca13f8
commit 7e2cefc3aa
3 changed files with 139 additions and 10 deletions

View File

@@ -23,7 +23,7 @@ import java.util.List;
*/
public class DefaultChatOptionsBuilder<T extends DefaultChatOptionsBuilder<T>> implements ChatOptions.Builder<T> {
private final DefaultChatOptions options = new DefaultChatOptions();
protected DefaultChatOptions options = new DefaultChatOptions();
protected T self() {
return (T) this;

View File

@@ -36,22 +36,28 @@ public class DefaultFunctionCallingOptionsBuilder
extends DefaultChatOptionsBuilder<DefaultFunctionCallingOptionsBuilder>
implements FunctionCallingOptions.Builder<DefaultFunctionCallingOptionsBuilder> {
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<FunctionCallback> 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<String> 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<String, Object> context) {
@@ -72,7 +78,7 @@ public class DefaultFunctionCallingOptionsBuilder
Map<String, Object> 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<String, Object> newContext = new HashMap<>(this.functionCallingOptions.getToolContext());
newContext.put(key, value);
this.functionCallingOptions.setToolContext(newContext);
return this;
return self();
}
public FunctionCallingOptions build() {

View File

@@ -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<String> 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<String, Object> 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);
}
}