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 e10487f14..f9769d927 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 @@ -17,6 +17,7 @@ package org.springframework.ai.vertexai.gemini; import java.util.ArrayList; +import java.util.HashMap; import java.util.HashSet; import java.util.List; import java.util.Map; @@ -28,6 +29,7 @@ import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; import com.fasterxml.jackson.annotation.JsonProperty; +import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel.ChatModel; @@ -68,7 +70,7 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions { /** * Optional. If specified, top k sampling will be used. */ - private @JsonProperty("topK") Float topK; + private @JsonProperty("topK") Integer topK; /** * Optional. The maximum number of tokens to generate. @@ -183,16 +185,11 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions { @Override public Integer getTopK() { - return (this.topK != null) ? this.topK.intValue() : null; + return this.topK; } - public void setTopK(Float topK) { - this.topK = topK; - } - - @JsonIgnore public void setTopK(Integer topK) { - this.topK = (topK != null) ? topK.floatValue() : null; + this.topK = topK; } public Integer getCandidateCount() { @@ -346,6 +343,67 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions { return fromOptions(this); } + public FunctionCallingOptions merge(ChatOptions options) { + VertexAiGeminiChatOptions.Builder builder = VertexAiGeminiChatOptions.builder(); + + // Merge chat-specific options + builder.model(options.getModel() != null ? options.getModel() : this.getModel()) + .maxOutputTokens(options.getMaxTokens() != null ? options.getMaxTokens() : this.getMaxOutputTokens()) + .stopSequences(options.getStopSequences() != null ? options.getStopSequences() : this.getStopSequences()) + .temperature(options.getTemperature() != null ? options.getTemperature() : this.getTemperature()) + .topP(options.getTopP() != null ? options.getTopP() : this.getTopP()) + .topK(options.getTopK() != null ? options.getTopK() : this.getTopK()); + + // Try to get function-specific properties if options is a FunctionCallingOptions + if (options instanceof FunctionCallingOptions functionOptions) { + builder.proxyToolCalls(functionOptions.getProxyToolCalls() != null ? functionOptions.getProxyToolCalls() + : this.proxyToolCalls); + + Set functions = new HashSet<>(); + if (this.functions != null) { + functions.addAll(this.functions); + } + if (functionOptions.getFunctions() != null) { + functions.addAll(functionOptions.getFunctions()); + } + builder.functions(functions); + + List functionCallbacks = new ArrayList<>(); + if (this.functionCallbacks != null) { + functionCallbacks.addAll(this.functionCallbacks); + } + if (functionOptions.getFunctionCallbacks() != null) { + functionCallbacks.addAll(functionOptions.getFunctionCallbacks()); + } + builder.functionCallbacks(functionCallbacks); + + Map context = new HashMap<>(); + if (this.toolContext != null) { + context.putAll(this.toolContext); + } + if (functionOptions.getToolContext() != null) { + context.putAll(functionOptions.getToolContext()); + } + builder.toolContext(context); + } + else { + // If not a FunctionCallingOptions, preserve current function-specific + // properties + builder.proxyToolCalls(this.proxyToolCalls); + builder.functions(this.functions != null ? new HashSet<>(this.functions) : null); + builder.functionCallbacks(this.functionCallbacks != null ? new ArrayList<>(this.functionCallbacks) : null); + builder.toolContext(this.toolContext != null ? new HashMap<>(this.toolContext) : null); + } + + // Preserve Vertex AI Gemini-specific properties + builder.candidateCount(this.candidateCount) + .responseMimeType(this.responseMimeType) + .googleSearchRetrieval(this.googleSearchRetrieval) + .safetySettings(this.safetySettings != null ? new ArrayList<>(this.safetySettings) : null); + + return builder.build(); + } + public enum TransportType { GRPC, REST @@ -371,7 +429,7 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions { return this; } - public Builder topK(Float topK) { + public Builder topK(Integer topK) { this.options.setTopK(topK); return this; } @@ -473,10 +531,10 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions { } /** - * @deprecated use {@link #topK(Float)} instead. + * @deprecated use {@link #topK(Integer)} instead. */ @Deprecated(forRemoval = true, since = "1.0.0-M5") - public Builder withTopK(Float topK) { + public Builder withTopK(Integer topK) { this.options.setTopK(topK); return this; } diff --git a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/CreateGeminiRequestTests.java b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/CreateGeminiRequestTests.java index bc3eea8a3..d92bf2b76 100644 --- a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/CreateGeminiRequestTests.java +++ b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/CreateGeminiRequestTests.java @@ -203,7 +203,7 @@ public class CreateGeminiRequestTests { .model("DEFAULT_MODEL") .temperature(66.6) .maxOutputTokens(100) - .topK(10.0f) + .topK(10) .topP(5.0) .stopSequences(List.of("stop1", "stop2")) .candidateCount(1) @@ -218,7 +218,7 @@ public class CreateGeminiRequestTests { assertThat(request.model().getModelName()).isEqualTo("DEFAULT_MODEL"); assertThat(request.model().getGenerationConfig().getTemperature()).isEqualTo(66.6f); assertThat(request.model().getGenerationConfig().getMaxOutputTokens()).isEqualTo(100); - assertThat(request.model().getGenerationConfig().getTopK()).isEqualTo(10.0f); + assertThat(request.model().getGenerationConfig().getTopK()).isEqualTo(10); assertThat(request.model().getGenerationConfig().getTopP()).isEqualTo(5.0f); assertThat(request.model().getGenerationConfig().getCandidateCount()).isEqualTo(1); assertThat(request.model().getGenerationConfig().getStopSequences(0)).isEqualTo("stop1"); 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 1758a2f46..6a10d1d27 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 @@ -17,6 +17,7 @@ package org.springframework.ai.zhipuai; import java.util.ArrayList; +import java.util.HashMap; import java.util.HashSet; import java.util.List; import java.util.Map; @@ -27,6 +28,7 @@ import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; import com.fasterxml.jackson.annotation.JsonProperty; +import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.zhipuai.api.ZhiPuAiApi; @@ -438,6 +440,67 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions { return fromOptions(this); } + public FunctionCallingOptions merge(ChatOptions options) { + ZhiPuAiChatOptions.Builder builder = ZhiPuAiChatOptions.builder(); + + // Merge chat-specific options + builder.model(options.getModel() != null ? options.getModel() : this.getModel()) + .maxTokens(options.getMaxTokens() != null ? options.getMaxTokens() : this.getMaxTokens()) + .stop(options.getStopSequences() != null ? options.getStopSequences() : this.getStopSequences()) + .temperature(options.getTemperature() != null ? options.getTemperature() : this.getTemperature()) + .topP(options.getTopP() != null ? options.getTopP() : this.getTopP()); + + // Try to get function-specific properties if options is a FunctionCallingOptions + if (options instanceof FunctionCallingOptions functionOptions) { + builder.proxyToolCalls(functionOptions.getProxyToolCalls() != null ? functionOptions.getProxyToolCalls() + : this.proxyToolCalls); + + Set functions = new HashSet<>(); + if (this.functions != null) { + functions.addAll(this.functions); + } + if (functionOptions.getFunctions() != null) { + functions.addAll(functionOptions.getFunctions()); + } + builder.functions(functions); + + List functionCallbacks = new ArrayList<>(); + if (this.functionCallbacks != null) { + functionCallbacks.addAll(this.functionCallbacks); + } + if (functionOptions.getFunctionCallbacks() != null) { + functionCallbacks.addAll(functionOptions.getFunctionCallbacks()); + } + builder.functionCallbacks(functionCallbacks); + + Map context = new HashMap<>(); + if (this.toolContext != null) { + context.putAll(this.toolContext); + } + if (functionOptions.getToolContext() != null) { + context.putAll(functionOptions.getToolContext()); + } + builder.toolContext(context); + } + else { + // If not a FunctionCallingOptions, preserve current function-specific + // properties + builder.proxyToolCalls(this.proxyToolCalls); + builder.functions(this.functions != null ? new HashSet<>(this.functions) : null); + builder.functionCallbacks(this.functionCallbacks != null ? new ArrayList<>(this.functionCallbacks) : null); + builder.toolContext(this.toolContext != null ? new HashMap<>(this.toolContext) : null); + } + + // Preserve ZhiPuAi-specific properties + builder.tools(this.tools) + .toolChoice(this.toolChoice) + .user(this.user) + .requestId(this.requestId) + .doSample(this.doSample); + + return builder.build(); + } + public static class Builder { protected ZhiPuAiChatOptions options; 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 8cc03c5ed..19cb98a3a 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 @@ -87,77 +87,77 @@ public interface ChatOptions extends ModelOptions { * Returns a copy of this {@link ChatOptions}. * @return a copy of this {@link ChatOptions} */ - ChatOptions copy(); + T copy(); /** * Creates a new {@link ChatOptions.Builder} to create the default * {@link ChatOptions}. * @return Returns a new {@link ChatOptions.Builder}. */ - static ChatOptions.Builder builder() { + static ChatOptions.Builder builder() { return new DefaultChatOptionsBuilder(); } /** * Builder for creating {@link ChatOptions} instance. */ - interface Builder> { + interface Builder { /** * Builds with the model to use for the chat. * @param model * @return the builder */ - B model(String model); + Builder model(String model); /** * Builds with the frequency penalty to use for the chat. * @param frequencyPenalty * @return the builder. */ - B frequencyPenalty(Double frequencyPenalty); + Builder frequencyPenalty(Double frequencyPenalty); /** * Builds with the maximum number of tokens to use for the chat. * @param maxTokens * @return the builder. */ - B maxTokens(Integer maxTokens); + Builder maxTokens(Integer maxTokens); /** * Builds with the presence penalty to use for the chat. * @param presencePenalty * @return the builder. */ - B presencePenalty(Double presencePenalty); + Builder presencePenalty(Double presencePenalty); /** * Builds with the stop sequences to use for the chat. * @param stopSequences * @return the builder. */ - B stopSequences(List stopSequences); + Builder stopSequences(List stopSequences); /** * Builds with the temperature to use for the chat. * @param temperature * @return the builder. */ - B temperature(Double temperature); + Builder temperature(Double temperature); /** * Builds with the top K to use for the chat. * @param topK * @return the builder. */ - B topK(Integer topK); + Builder topK(Integer topK); /** * Builds with the top P to use for the chat. * @param topP * @return the builder. */ - B topP(Double topP); + Builder topP(Double topP); /** * Build the {@link ChatOptions}. diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/DefaultChatOptions.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/DefaultChatOptions.java index 9feca34a9..1af33bf34 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/DefaultChatOptions.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/DefaultChatOptions.java @@ -16,6 +16,8 @@ package org.springframework.ai.chat.prompt; +import java.util.ArrayList; +import java.util.Collections; import java.util.List; /** @@ -77,7 +79,7 @@ public class DefaultChatOptions implements ChatOptions { @Override public List getStopSequences() { - return this.stopSequences; + return this.stopSequences != null ? Collections.unmodifiableList(this.stopSequences) : null; } public void setStopSequences(List stopSequences) { @@ -112,17 +114,18 @@ public class DefaultChatOptions implements ChatOptions { } @Override - public ChatOptions copy() { - return ChatOptions.builder() - .model(this.model) - .frequencyPenalty(this.frequencyPenalty) - .maxTokens(this.maxTokens) - .presencePenalty(this.presencePenalty) - .stopSequences(this.stopSequences != null ? List.copyOf(this.stopSequences) : null) - .temperature(this.temperature) - .topK(this.topK) - .topP(this.topP) - .build(); + @SuppressWarnings("unchecked") + public T copy() { + DefaultChatOptions copy = new DefaultChatOptions(); + copy.setModel(this.getModel()); + copy.setFrequencyPenalty(this.getFrequencyPenalty()); + copy.setMaxTokens(this.getMaxTokens()); + copy.setPresencePenalty(this.getPresencePenalty()); + copy.setStopSequences(this.getStopSequences() != null ? new ArrayList<>(this.getStopSequences()) : null); + copy.setTemperature(this.getTemperature()); + copy.setTopK(this.getTopK()); + copy.setTopP(this.getTopP()); + return (T) copy; } } 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 e10f796a7..47ba58401 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 @@ -21,7 +21,7 @@ import java.util.List; /** * Implementation of {@link ChatOptions.Builder} to create {@link DefaultChatOptions}. */ -public class DefaultChatOptionsBuilder> implements ChatOptions.Builder { +public class DefaultChatOptionsBuilder implements ChatOptions.Builder { protected DefaultChatOptions options; @@ -33,52 +33,48 @@ public class DefaultChatOptionsBuilder> i this.options = options; } - protected T self() { - return (T) this; - } - - public T model(String model) { + public DefaultChatOptionsBuilder model(String model) { this.options.setModel(model); - return self(); + return this; } - public T frequencyPenalty(Double frequencyPenalty) { + public DefaultChatOptionsBuilder frequencyPenalty(Double frequencyPenalty) { this.options.setFrequencyPenalty(frequencyPenalty); - return self(); + return this; } - public T maxTokens(Integer maxTokens) { + public DefaultChatOptionsBuilder maxTokens(Integer maxTokens) { this.options.setMaxTokens(maxTokens); - return self(); + return this; } - public T presencePenalty(Double presencePenalty) { + public DefaultChatOptionsBuilder presencePenalty(Double presencePenalty) { this.options.setPresencePenalty(presencePenalty); - return self(); + return this; } - public T stopSequences(List stop) { + public DefaultChatOptionsBuilder stopSequences(List stop) { this.options.setStopSequences(stop); - return self(); + return this; } - public T temperature(Double temperature) { + public DefaultChatOptionsBuilder temperature(Double temperature) { this.options.setTemperature(temperature); - return self(); + return this; } - public T topK(Integer topK) { + public DefaultChatOptionsBuilder topK(Integer topK) { this.options.setTopK(topK); - return self(); + return this; } - public T topP(Double topP) { + public DefaultChatOptionsBuilder topP(Double topP) { this.options.setTopP(topP); - return self(); + return this; } public ChatOptions build() { - return this.options; + return this.options.copy(); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/DefaultFunctionCallingOptions.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/DefaultFunctionCallingOptions.java index 180f8ddfb..14ccccf10 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/DefaultFunctionCallingOptions.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/DefaultFunctionCallingOptions.java @@ -25,7 +25,6 @@ import java.util.Map; import java.util.Set; import org.springframework.ai.chat.prompt.ChatOptions; -import org.springframework.ai.chat.prompt.DefaultChatOptions; import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; import org.springframework.util.StringUtils; @@ -37,7 +36,7 @@ import org.springframework.util.StringUtils; * @author Thomas Vitale * @author Ilayaperumal Gopinathan */ -public class DefaultFunctionCallingOptions extends DefaultChatOptions implements FunctionCallingOptions { +public class DefaultFunctionCallingOptions implements FunctionCallingOptions { private List functionCallbacks = new ArrayList<>(); @@ -47,6 +46,22 @@ public class DefaultFunctionCallingOptions extends DefaultChatOptions implements private Map context = new HashMap<>(); + private String model; + + private Double frequencyPenalty; + + private Integer maxTokens; + + private Double presencePenalty; + + private List stopSequences; + + private Double temperature; + + private Integer topK; + + private Double topP; + @Override public List getFunctionCallbacks() { return Collections.unmodifiableList(this.functionCallbacks); @@ -86,73 +101,99 @@ public class DefaultFunctionCallingOptions extends DefaultChatOptions implements } @Override - public FunctionCallingOptions copy() { - return FunctionCallingOptions.builder() - .model(this.getModel()) - .frequencyPenalty(this.getFrequencyPenalty()) - .maxTokens(this.getMaxTokens()) - .presencePenalty(this.getPresencePenalty()) - .stopSequences(this.getStopSequences() != null ? new ArrayList<>(this.getStopSequences()) : null) - .temperature(this.getTemperature()) - .topK(this.getTopK()) - .topP(this.getTopP()) - .functions(new HashSet<>(this.functions)) - .functionCallbacks(new ArrayList<>(this.functionCallbacks)) - .proxyToolCalls(this.proxyToolCalls) - .toolContext(new HashMap<>(this.getToolContext())) - .build(); + public String getModel() { + return this.model; } - public FunctionCallingOptions merge(FunctionCallingOptions options) { + public void setModel(String model) { + this.model = model; + } - var builder = FunctionCallingOptions.builder() - .model(StringUtils.hasText(options.getModel()) ? options.getModel() : this.getModel()) - .frequencyPenalty( - options.getFrequencyPenalty() != null ? options.getFrequencyPenalty() : this.getFrequencyPenalty()) - .maxTokens(options.getMaxTokens() != null ? options.getMaxTokens() : this.getMaxTokens()) - .presencePenalty( - options.getPresencePenalty() != null ? options.getPresencePenalty() : this.getPresencePenalty()) - .stopSequences(options.getStopSequences() != null ? options.getStopSequences() : this.getStopSequences()) - .temperature(options.getTemperature() != null ? options.getTemperature() : this.getTemperature()) - .topK(options.getTopK() != null ? options.getTopK() : this.getTopK()) - .topP(options.getTopP() != null ? options.getTopP() : this.getTopP()); + @Override + public Double getFrequencyPenalty() { + return this.frequencyPenalty; + } - builder.proxyToolCalls(options.getProxyToolCalls() != null ? options.getProxyToolCalls() : this.proxyToolCalls); + public void setFrequencyPenalty(Double frequencyPenalty) { + this.frequencyPenalty = frequencyPenalty; + } - Set functions = new HashSet<>(); - if (!CollectionUtils.isEmpty(this.functions)) { - functions.addAll(this.functions); - } - if (!CollectionUtils.isEmpty(options.getFunctions())) { - functions.addAll(options.getFunctions()); - } - builder.functions(functions); + @Override + public Integer getMaxTokens() { + return this.maxTokens; + } - List functionCallbacks = new ArrayList<>(); - if (!CollectionUtils.isEmpty(this.functionCallbacks)) { - functionCallbacks.addAll(this.functionCallbacks); - } - if (!CollectionUtils.isEmpty(options.getFunctionCallbacks())) { - functionCallbacks.addAll(options.getFunctionCallbacks()); - } - builder.functionCallbacks(functionCallbacks); + public void setMaxTokens(Integer maxTokens) { + this.maxTokens = maxTokens; + } - Map context = new HashMap<>(); - if (!CollectionUtils.isEmpty(this.context)) { - context.putAll(this.context); - } - if (!CollectionUtils.isEmpty(options.getToolContext())) { - context.putAll(options.getToolContext()); - } - builder.toolContext(context); + @Override + public Double getPresencePenalty() { + return this.presencePenalty; + } - return builder.build(); + public void setPresencePenalty(Double presencePenalty) { + this.presencePenalty = presencePenalty; + } + + @Override + public List getStopSequences() { + return this.stopSequences != null ? Collections.unmodifiableList(this.stopSequences) : null; + } + + public void setStopSequences(List stopSequences) { + this.stopSequences = stopSequences; + } + + @Override + public Double getTemperature() { + return this.temperature; + } + + public void setTemperature(Double temperature) { + this.temperature = temperature; + } + + @Override + public Integer getTopK() { + return this.topK; + } + + public void setTopK(Integer topK) { + this.topK = topK; + } + + @Override + public Double getTopP() { + return this.topP; + } + + public void setTopP(Double topP) { + this.topP = topP; + } + + @Override + @SuppressWarnings("unchecked") + public T copy() { + DefaultFunctionCallingOptions copy = new DefaultFunctionCallingOptions(); + copy.setModel(this.getModel()); + copy.setFrequencyPenalty(this.getFrequencyPenalty()); + copy.setMaxTokens(this.getMaxTokens()); + copy.setPresencePenalty(this.getPresencePenalty()); + copy.setStopSequences(this.getStopSequences() != null ? new ArrayList<>(this.getStopSequences()) : null); + copy.setTemperature(this.getTemperature()); + copy.setTopK(this.getTopK()); + copy.setTopP(this.getTopP()); + copy.setFunctions(new HashSet<>(this.functions)); + copy.setFunctionCallbacks(new ArrayList<>(this.functionCallbacks)); + copy.setProxyToolCalls(this.proxyToolCalls); + copy.setToolContext(new HashMap<>(this.getToolContext())); + return (T) copy; } public FunctionCallingOptions merge(ChatOptions options) { - - var builder = FunctionCallingOptions.builder() - .model(StringUtils.hasText(options.getModel()) ? options.getModel() : this.getModel()) + FunctionCallingOptions.Builder builder = FunctionCallingOptions.builder(); + builder.model(StringUtils.hasText(options.getModel()) ? options.getModel() : this.getModel()) .frequencyPenalty( options.getFrequencyPenalty() != null ? options.getFrequencyPenalty() : this.getFrequencyPenalty()) .maxTokens(options.getMaxTokens() != null ? options.getMaxTokens() : this.getMaxTokens()) @@ -163,6 +204,47 @@ public class DefaultFunctionCallingOptions extends DefaultChatOptions implements .topK(options.getTopK() != null ? options.getTopK() : this.getTopK()) .topP(options.getTopP() != null ? options.getTopP() : this.getTopP()); + // Try to get function-specific properties if options is a FunctionCallingOptions + if (options instanceof FunctionCallingOptions functionOptions) { + builder.proxyToolCalls(functionOptions.getProxyToolCalls() != null ? functionOptions.getProxyToolCalls() + : this.proxyToolCalls); + + Set functions = new HashSet<>(); + if (!CollectionUtils.isEmpty(this.functions)) { + functions.addAll(this.functions); + } + if (!CollectionUtils.isEmpty(functionOptions.getFunctions())) { + functions.addAll(functionOptions.getFunctions()); + } + builder.functions(functions); + + List functionCallbacks = new ArrayList<>(); + if (!CollectionUtils.isEmpty(this.functionCallbacks)) { + functionCallbacks.addAll(this.functionCallbacks); + } + if (!CollectionUtils.isEmpty(functionOptions.getFunctionCallbacks())) { + functionCallbacks.addAll(functionOptions.getFunctionCallbacks()); + } + builder.functionCallbacks(functionCallbacks); + + Map context = new HashMap<>(); + if (!CollectionUtils.isEmpty(this.context)) { + context.putAll(this.context); + } + if (!CollectionUtils.isEmpty(functionOptions.getToolContext())) { + context.putAll(functionOptions.getToolContext()); + } + builder.toolContext(context); + } + else { + // If not a FunctionCallingOptions, preserve current function-specific + // properties + builder.proxyToolCalls(this.proxyToolCalls); + builder.functions(new HashSet<>(this.functions)); + builder.functionCallbacks(new ArrayList<>(this.functionCallbacks)); + builder.toolContext(new HashMap<>(this.context)); + } + return builder.build(); } 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 13cf89539..15f8926e7 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 @@ -22,7 +22,6 @@ import java.util.List; import java.util.Map; import java.util.Set; -import org.springframework.ai.chat.prompt.DefaultChatOptionsBuilder; import org.springframework.util.Assert; /** @@ -32,63 +31,120 @@ import org.springframework.util.Assert; * @author Thomas Vitale * @author Ilayaperumal Gopinathan */ -public class DefaultFunctionCallingOptionsBuilder - extends DefaultChatOptionsBuilder - implements FunctionCallingOptions.Builder { +public class DefaultFunctionCallingOptionsBuilder implements FunctionCallingOptions.Builder { + + private final DefaultFunctionCallingOptions options; public DefaultFunctionCallingOptionsBuilder() { - // Set the options in the parent class to be the same instance - super(new DefaultFunctionCallingOptions()); + this.options = new DefaultFunctionCallingOptions(); } - public DefaultFunctionCallingOptionsBuilder functionCallbacks(List functionCallbacks) { - ((FunctionCallingOptions) this.options).setFunctionCallbacks(functionCallbacks); - return self(); + // Function calling specific methods + @Override + public FunctionCallingOptions.Builder functionCallbacks(List functionCallbacks) { + this.options.setFunctionCallbacks(functionCallbacks); + return this; } - public DefaultFunctionCallingOptionsBuilder functionCallbacks(FunctionCallback... functionCallbacks) { + @Override + public FunctionCallingOptions.Builder functionCallbacks(FunctionCallback... functionCallbacks) { Assert.notNull(functionCallbacks, "FunctionCallbacks must not be null"); - ((FunctionCallingOptions) this.options).setFunctionCallbacks(List.of(functionCallbacks)); - return self(); + this.options.setFunctionCallbacks(List.of(functionCallbacks)); + return this; } - public DefaultFunctionCallingOptionsBuilder functions(Set functions) { - ((FunctionCallingOptions) this.options).setFunctions(functions); - return self(); + @Override + public FunctionCallingOptions.Builder functions(Set functions) { + this.options.setFunctions(functions); + return this; } - public DefaultFunctionCallingOptionsBuilder function(String function) { + @Override + public FunctionCallingOptions.Builder function(String function) { Assert.notNull(function, "Function must not be null"); - var set = new HashSet<>(((FunctionCallingOptions) this.options).getFunctions()); + var set = new HashSet<>(this.options.getFunctions()); set.add(function); - ((FunctionCallingOptions) this.options).setFunctions(set); - return self(); + this.options.setFunctions(set); + return this; } - public DefaultFunctionCallingOptionsBuilder proxyToolCalls(Boolean proxyToolCalls) { - ((FunctionCallingOptions) this.options).setProxyToolCalls(proxyToolCalls); - return self(); + @Override + public FunctionCallingOptions.Builder proxyToolCalls(Boolean proxyToolCalls) { + this.options.setProxyToolCalls(proxyToolCalls); + return this; } - public DefaultFunctionCallingOptionsBuilder toolContext(Map context) { + @Override + public FunctionCallingOptions.Builder toolContext(Map context) { Assert.notNull(context, "Tool context must not be null"); - Map newContext = new HashMap<>(((FunctionCallingOptions) this.options).getToolContext()); + Map newContext = new HashMap<>(this.options.getToolContext()); newContext.putAll(context); - ((FunctionCallingOptions) this.options).setToolContext(newContext); - return self(); + this.options.setToolContext(newContext); + return this; } - public DefaultFunctionCallingOptionsBuilder toolContext(String key, Object value) { + @Override + public FunctionCallingOptions.Builder toolContext(String key, Object value) { Assert.notNull(key, "Key must not be null"); Assert.notNull(value, "Value must not be null"); - Map newContext = new HashMap<>(((FunctionCallingOptions) this.options).getToolContext()); + Map newContext = new HashMap<>(this.options.getToolContext()); newContext.put(key, value); - ((FunctionCallingOptions) this.options).setToolContext(newContext); - return self(); + this.options.setToolContext(newContext); + return this; } + // ChatOptions.Builder methods with covariant return type + @Override + public FunctionCallingOptions.Builder model(String model) { + this.options.setModel(model); + return this; + } + + @Override + public FunctionCallingOptions.Builder frequencyPenalty(Double frequencyPenalty) { + this.options.setFrequencyPenalty(frequencyPenalty); + return this; + } + + @Override + public FunctionCallingOptions.Builder maxTokens(Integer maxTokens) { + this.options.setMaxTokens(maxTokens); + return this; + } + + @Override + public FunctionCallingOptions.Builder presencePenalty(Double presencePenalty) { + this.options.setPresencePenalty(presencePenalty); + return this; + } + + @Override + public FunctionCallingOptions.Builder stopSequences(List stop) { + this.options.setStopSequences(stop); + return this; + } + + @Override + public FunctionCallingOptions.Builder temperature(Double temperature) { + this.options.setTemperature(temperature); + return this; + } + + @Override + public FunctionCallingOptions.Builder topK(Integer topK) { + this.options.setTopK(topK); + return this; + } + + @Override + public FunctionCallingOptions.Builder topP(Double topP) { + this.options.setTopP(topP); + return this; + } + + @Override public FunctionCallingOptions build() { - return ((FunctionCallingOptions) this.options); + return this.options.copy(); } } 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 0859df140..ba8529e68 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 @@ -35,7 +35,7 @@ public interface FunctionCallingOptions extends ChatOptions { * @return Returns {@link DefaultFunctionCallingOptionsBuilder} to create a new * instance of {@link FunctionCallingOptions}. */ - static FunctionCallingOptions.Builder builder() { + static FunctionCallingOptions.Builder builder() { return new DefaultFunctionCallingOptionsBuilder(); } @@ -87,49 +87,49 @@ public interface FunctionCallingOptions extends ChatOptions { /** * Builder for creating {@link FunctionCallingOptions} instance. */ - interface Builder> extends ChatOptions.Builder { + interface Builder extends ChatOptions.Builder { /** * The list of Function Callbacks to be registered with the Chat model. * @param functionCallbacks the list of Function Callbacks. * @return the FunctionCallOptions Builder. */ - T functionCallbacks(List functionCallbacks); + Builder functionCallbacks(List functionCallbacks); /** * The Function Callbacks to be registered with the Chat model. * @param functionCallbacks the function callbacks. * @return the FunctionCallOptions Builder. */ - T functionCallbacks(FunctionCallback... functionCallbacks); + Builder functionCallbacks(FunctionCallback... functionCallbacks); /** * {@link Set} of function names to be registered with the Chat model. * @param functions the {@link Set} of function names * @return the FunctionCallOptions Builder. */ - T functions(Set functions); + Builder functions(Set functions); /** * The function name to be registered with the chat model. * @param function the name of the function. * @return the FunctionCallOptions Builder. */ - T function(String function); + Builder function(String function); /** * Boolean flag to indicate if the proxy ToolCalls is enabled. * @param proxyToolCalls boolean value to enable proxy ToolCalls. * @return the FunctionCallOptions Builder. */ - T proxyToolCalls(Boolean proxyToolCalls); + Builder proxyToolCalls(Boolean proxyToolCalls); /** * Add a {@link Map} of context values into tool context. * @param context the map representing the tool context. * @return the FunctionCallOptions Builder. */ - T toolContext(Map context); + Builder toolContext(Map context); /** * Add a specific key/value pair to the tool context. @@ -137,14 +137,41 @@ public interface FunctionCallingOptions extends ChatOptions { * @param value the corresponding value. * @return the FunctionCallOptions Builder. */ - T toolContext(String key, Object value); + Builder toolContext(String key, Object value); /** * Builds the {@link FunctionCallingOptions}. * @return the FunctionCalling options. */ + @Override FunctionCallingOptions build(); + // Override all ChatOptions.Builder methods to return + // FunctionCallingOptions.Builder + @Override + Builder model(String model); + + @Override + Builder frequencyPenalty(Double frequencyPenalty); + + @Override + Builder maxTokens(Integer maxTokens); + + @Override + Builder presencePenalty(Double presencePenalty); + + @Override + Builder stopSequences(List stopSequences); + + @Override + Builder temperature(Double temperature); + + @Override + Builder topK(Integer topK); + + @Override + Builder topP(Double topP); + } } 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 deleted file mode 100644 index fd15f8465..000000000 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatBuilderTests.java +++ /dev/null @@ -1,108 +0,0 @@ -/* - * 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.chat; - -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.Prompt; -import org.springframework.ai.model.function.FunctionCallback; -import org.springframework.ai.model.function.FunctionCallingOptions; - -import static org.assertj.core.api.Assertions.assertThat; - -/** - * Unit Tests for {@link Prompt}. - * - * @author youngmon - * @since 0.8.1 - */ -public class ChatBuilderTests { - - @Test - void createNewChatOptionsTest() { - Double temperature = 1.1; - Double topP = 2.2; - Integer topK = 111; - - ChatOptions options = ChatOptions.builder().temperature(temperature).topK(topK).topP(topP).build(); - - assertThat(options.getTemperature()).isEqualTo(temperature); - assertThat(options.getTopP()).isEqualTo(topP); - assertThat(options.getTopK()).isEqualTo(topK); - } - - @Test - void duplicateChatOptionsTest() { - Double initTemperature = 1.1; - Double initTopP = 2.2; - Integer initTopK = 111; - - ChatOptions options1 = ChatOptions.builder().temperature(initTemperature).topP(initTopP).topK(initTopK).build(); - - ChatOptions options2 = options1.copy(); - - assertThat(options2.getTemperature()).isEqualTo(initTemperature); - assertThat(options2.getTopP()).isEqualTo(initTopP); - assertThat(options2.getTopK()).isEqualTo(initTopK); - - } - - @Test - void createFunctionCallingOptionTest() { - Double temperature = 1.1; - Double topP = 2.2; - Integer topK = 111; - List functionCallbacks = new ArrayList<>(); - Set functions = new HashSet<>(); - - String func = "func"; - FunctionCallback cb = FunctionCallback.builder() - .function("cb", i -> i) - .description("cb") - .inputType(Integer.class) - .build(); - - functions.add(func); - functionCallbacks.add(cb); - - FunctionCallingOptions options = FunctionCallingOptions.builder() - .functionCallbacks(functionCallbacks) - .functions(functions) - .topK(topK) - .topP(topP) - .temperature(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); - - } - -} diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/prompt/ChatOptionsBuilderTests.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/prompt/ChatOptionsBuilderTests.java new file mode 100644 index 000000000..0a689d026 --- /dev/null +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/prompt/ChatOptionsBuilderTests.java @@ -0,0 +1,178 @@ +/* + * 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.chat.prompt; + +import java.util.ArrayList; +import java.util.HashSet; +import java.util.List; +import java.util.Set; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import org.springframework.ai.model.function.FunctionCallback; +import org.springframework.ai.model.function.FunctionCallingOptions; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.AssertionsForClassTypes.assertThatThrownBy; + +/** + * Unit Tests for {@link ChatOptions} builder. + * + * @author youngmon + * @author Mark Pollack + * @since 1.0.0 + */ +public class ChatOptionsBuilderTests { + + private ChatOptions.Builder builder; + + @BeforeEach + void setUp() { + builder = ChatOptions.builder(); + } + + @Test + void shouldBuildWithAllOptions() { + ChatOptions options = builder.model("gpt-4") + .maxTokens(100) + .temperature(0.7) + .topP(1.0) + .topK(40) + .stopSequences(List.of("stop1", "stop2")) + .build(); + + assertThat(options.getModel()).isEqualTo("gpt-4"); + assertThat(options.getMaxTokens()).isEqualTo(100); + assertThat(options.getTemperature()).isEqualTo(0.7); + assertThat(options.getTopP()).isEqualTo(1.0); + assertThat(options.getTopK()).isEqualTo(40); + assertThat(options.getStopSequences()).containsExactly("stop1", "stop2"); + } + + @Test + void shouldBuildWithMinimalOptions() { + ChatOptions options = builder.model("gpt-4").build(); + + assertThat(options.getModel()).isEqualTo("gpt-4"); + assertThat(options.getMaxTokens()).isNull(); + assertThat(options.getTemperature()).isNull(); + assertThat(options.getTopP()).isNull(); + assertThat(options.getTopK()).isNull(); + assertThat(options.getStopSequences()).isNull(); + } + + @Test + void shouldCopyOptions() { + ChatOptions original = builder.model("gpt-4") + .maxTokens(100) + .temperature(0.7) + .topP(1.0) + .topK(40) + .stopSequences(List.of("stop1", "stop2")) + .build(); + + ChatOptions copy = original.copy(); + + // Then + assertThat(copy).usingRecursiveComparison().isEqualTo(original); + // Verify collections are actually copied + assertThat(copy.getStopSequences()).isNotSameAs(original.getStopSequences()); + } + + @Test + void shouldUpcastToChatOptions() { + // Given + FunctionCallback callback = FunctionCallback.builder() + .function("function1", x -> "result") + .description("Test function") + .inputType(String.class) + .build(); + + FunctionCallingOptions functionOptions = FunctionCallingOptions.builder() + .model("gpt-4") + .maxTokens(100) + .temperature(0.7) + .topP(1.0) + .topK(40) + .stopSequences(List.of("stop1", "stop2")) + .functions(Set.of("function1", "function2")) + .functionCallbacks(List.of(callback)) + .build(); + + // When + ChatOptions chatOptions = functionOptions; + + // Then + assertThat(chatOptions.getModel()).isEqualTo("gpt-4"); + assertThat(chatOptions.getMaxTokens()).isEqualTo(100); + assertThat(chatOptions.getTemperature()).isEqualTo(0.7); + assertThat(chatOptions.getTopP()).isEqualTo(1.0); + assertThat(chatOptions.getTopK()).isEqualTo(40); + assertThat(chatOptions.getStopSequences()).containsExactly("stop1", "stop2"); + } + + @Test + void shouldAllowBuilderReuse() { + // When + ChatOptions options1 = builder.model("model1").temperature(0.7).build(); + ChatOptions options2 = builder.model("model2").build(); + + // Then + assertThat(options1.getModel()).isEqualTo("model1"); + assertThat(options1.getTemperature()).isEqualTo(0.7); + assertThat(options2.getModel()).isEqualTo("model2"); + assertThat(options2.getTemperature()).isEqualTo(0.7); // Retains previous value + } + + @Test + void shouldReturnSameBuilderInstanceOnEachMethod() { + // When + ChatOptions.Builder returnedBuilder = builder.model("test"); + + // Then + assertThat(returnedBuilder).isSameAs(builder); + } + + @Test + void shouldHaveExpectedDefaultValues() { + // When + ChatOptions options = builder.build(); + + // Then + assertThat(options.getModel()).isNull(); + assertThat(options.getTemperature()).isNull(); + assertThat(options.getMaxTokens()).isNull(); + assertThat(options.getTopP()).isNull(); + assertThat(options.getTopK()).isNull(); + assertThat(options.getFrequencyPenalty()).isNull(); + assertThat(options.getPresencePenalty()).isNull(); + assertThat(options.getStopSequences()).isNull(); + } + + @Test + void shouldBeImmutableAfterBuild() { + // Given + List stopSequences = new ArrayList<>(List.of("stop1", "stop2")); + ChatOptions options = builder.stopSequences(stopSequences).build(); + + // Then + assertThatThrownBy(() -> options.getStopSequences().add("stop3")) + .isInstanceOf(UnsupportedOperationException.class); + } + +} 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 7cc828cd7..ba40504d8 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 @@ -19,6 +19,7 @@ package org.springframework.ai.model.function; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.BeforeEach; +import java.util.ArrayList; import java.util.HashMap; import java.util.HashSet; import java.util.List; @@ -331,4 +332,246 @@ class DefaultFunctionCallingOptionsBuilderTests { assertThat(chatOptions.getTopP()).isEqualTo(0.9); } + @Test + void shouldBuildWithEmptyFunctionCallbacks() { + // When + FunctionCallingOptions options = builder.functionCallbacks(List.of()).build(); + + // Then + assertThat(options.getFunctionCallbacks()).isEmpty(); + } + + @Test + void shouldBuildWithEmptyFunctions() { + // When + FunctionCallingOptions options = builder.functions(Set.of()).build(); + + // Then + assertThat(options.getFunctions()).isEmpty(); + } + + @Test + void shouldBuildWithEmptyToolContext() { + // When + FunctionCallingOptions options = builder.toolContext(Map.of()).build(); + + // Then + assertThat(options.getToolContext()).isEmpty(); + } + + @Test + void shouldDeduplicateFunctions() { + // When + FunctionCallingOptions options = builder.function("function1") + .function("function1") // Duplicate + .function("function2") + .build(); + + // Then + assertThat(options.getFunctions()).hasSize(2).containsExactlyInAnyOrder("function1", "function2"); + } + + @Test + void shouldCopyAllOptions() { + // Given + FunctionCallback callback = FunctionCallback.builder() + .function("test", (String input) -> "result") + .description("Test function") + .inputType(String.class) + .build(); + FunctionCallingOptions original = 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) + .function("function1") + .proxyToolCalls(true) + .toolContext("key1", "value1") + .build(); + + // When + FunctionCallingOptions copy = original.copy(); + + // Then + assertThat(copy).usingRecursiveComparison().isEqualTo(original); + // Verify collections are actually copied + assertThat(copy.getFunctionCallbacks()).isNotSameAs(original.getFunctionCallbacks()); + assertThat(copy.getFunctions()).isNotSameAs(original.getFunctions()); + assertThat(copy.getToolContext()).isNotSameAs(original.getToolContext()); + } + + @Test + void shouldMergeWithFunctionCallingOptions() { + // Given + FunctionCallback callback1 = FunctionCallback.builder() + .function("test1", (String input) -> "result1") + .description("Test function 1") + .inputType(String.class) + .build(); + FunctionCallback callback2 = FunctionCallback.builder() + .function("test2", (String input) -> "result2") + .description("Test function 2") + .inputType(String.class) + .build(); + + DefaultFunctionCallingOptions options1 = (DefaultFunctionCallingOptions) builder.model("gpt-4") + .temperature(0.8) + .functionCallbacks(callback1) + .function("function1") + .proxyToolCalls(true) + .toolContext("key1", "value1") + .build(); + + DefaultFunctionCallingOptions options2 = (DefaultFunctionCallingOptions) FunctionCallingOptions.builder() + .model("gpt-3.5") + .maxTokens(100) + .functionCallbacks(callback2) + .function("function2") + .proxyToolCalls(false) + .toolContext("key2", "value2") + .build(); + + // When + FunctionCallingOptions merged = options1.merge(options2); + + // Then + assertThat(merged.getModel()).isEqualTo("gpt-3.5"); // Overridden + assertThat(merged.getTemperature()).isEqualTo(0.8); // Kept + assertThat(merged.getMaxTokens()).isEqualTo(100); // Added + assertThat(merged.getFunctionCallbacks()).containsExactly(callback1, callback2); // Combined + assertThat(merged.getFunctions()).containsExactlyInAnyOrder("function1", "function2"); // Combined + assertThat(merged.getProxyToolCalls()).isFalse(); // Overridden + assertThat(merged.getToolContext()).containsEntry("key1", "value1").containsEntry("key2", "value2"); // Combined + } + + @Test + void shouldMergeWithChatOptions() { + // Given + FunctionCallback callback = FunctionCallback.builder() + .function("test", (String input) -> "result") + .description("Test function") + .inputType(String.class) + .build(); + + DefaultFunctionCallingOptions options1 = (DefaultFunctionCallingOptions) builder.model("gpt-4") + .temperature(0.8) + .functionCallbacks(callback) + .function("function1") + .proxyToolCalls(true) + .toolContext("key1", "value1") + .build(); + + ChatOptions options2 = ChatOptions.builder().model("gpt-3.5").maxTokens(100).build(); + + // When + FunctionCallingOptions merged = options1.merge(options2); + + // Then + assertThat(merged.getModel()).isEqualTo("gpt-3.5"); // Overridden + assertThat(merged.getTemperature()).isEqualTo(0.8); // Kept + assertThat(merged.getMaxTokens()).isEqualTo(100); // Added + // Function-specific options should be preserved + assertThat(merged.getFunctionCallbacks()).containsExactly(callback); + assertThat(merged.getFunctions()).containsExactly("function1"); + assertThat(merged.getProxyToolCalls()).isTrue(); + assertThat(merged.getToolContext()).containsEntry("key1", "value1"); + } + + @Test + void shouldAllowBuilderReuse() { + // Given + FunctionCallback callback1 = FunctionCallback.builder() + .function("test1", (String input) -> "result1") + .description("Test function 1") + .inputType(String.class) + .build(); + FunctionCallback callback2 = FunctionCallback.builder() + .function("test2", (String input) -> "result2") + .description("Test function 2") + .inputType(String.class) + .build(); + + // When + FunctionCallingOptions options1 = builder.model("model1").temperature(0.7).functionCallbacks(callback1).build(); + + FunctionCallingOptions options2 = builder.model("model2").functionCallbacks(callback2).build(); + + // Then + assertThat(options1.getModel()).isEqualTo("model1"); + assertThat(options1.getTemperature()).isEqualTo(0.7); + assertThat(options1.getFunctionCallbacks()).containsExactly(callback1); + + assertThat(options2.getModel()).isEqualTo("model2"); + assertThat(options2.getTemperature()).isEqualTo(0.7); // Retains previous value + assertThat(options2.getFunctionCallbacks()).containsExactly(callback2); // Replaces + // previous + // callbacks + } + + @Test + void shouldReturnSameBuilderInstanceOnEachMethod() { + // When + FunctionCallingOptions.Builder returnedBuilder = builder.model("test"); + + // Then + assertThat(returnedBuilder).isSameAs(builder); + } + + @Test + void shouldHaveExpectedDefaultValues() { + // When + FunctionCallingOptions options = builder.build(); + + // Then + // ChatOptions defaults + assertThat(options.getModel()).isNull(); + assertThat(options.getTemperature()).isNull(); + assertThat(options.getMaxTokens()).isNull(); + assertThat(options.getTopP()).isNull(); + assertThat(options.getTopK()).isNull(); + assertThat(options.getFrequencyPenalty()).isNull(); + assertThat(options.getPresencePenalty()).isNull(); + assertThat(options.getStopSequences()).isNull(); + + // FunctionCallingOptions specific defaults + assertThat(options.getFunctionCallbacks()).isEmpty(); + assertThat(options.getFunctions()).isEmpty(); + assertThat(options.getToolContext()).isEmpty(); + assertThat(options.getProxyToolCalls()).isFalse(); + } + + @Test + void shouldBeImmutableAfterBuild() { + // Given + FunctionCallback callback = FunctionCallback.builder() + .function("test", (String input) -> "result") + .description("Test function") + .inputType(String.class) + .build(); + + List stopSequences = new ArrayList<>(List.of("stop1", "stop2")); + Set functions = new HashSet<>(Set.of("function1", "function2")); + Map context = new HashMap<>(Map.of("key1", "value1")); + + FunctionCallingOptions options = builder.stopSequences(stopSequences) + .functionCallbacks(callback) + .functions(functions) + .toolContext(context) + .build(); + + // Then + assertThatThrownBy(() -> options.getStopSequences().add("stop3")) + .isInstanceOf(UnsupportedOperationException.class); + assertThatThrownBy(() -> options.getFunctionCallbacks().add(callback)) + .isInstanceOf(UnsupportedOperationException.class); + assertThatThrownBy(() -> options.getFunctions().add("function3")) + .isInstanceOf(UnsupportedOperationException.class); + assertThatThrownBy(() -> options.getToolContext().put("key2", "value2")) + .isInstanceOf(UnsupportedOperationException.class); + } + }