Support reasoning_effort in AzureOpenAiChatOptions

* Adds reasoningEffort field to AzureOpenAiChatOptions builder, copy, equals, hashCode
* Propagates value to ChatCompletionsOptions

Fixes #2703

Signed-off-by: Andres da Silva Santos <40636137+andresssantos@users.noreply.github.com>
This commit is contained in:
Andres da Silva Santos
2025-05-08 09:53:29 -03:00
committed by Ilayaperumal Gopinathan
parent 69c081c778
commit fb402a6718
3 changed files with 60 additions and 11 deletions

View File

@@ -16,8 +16,6 @@
package org.springframework.ai.azure.openai; package org.springframework.ai.azure.openai;
import com.azure.ai.openai.models.ChatCompletionsJsonSchemaResponseFormat;
import com.azure.ai.openai.models.ChatCompletionsJsonSchemaResponseFormatJsonSchema;
import java.util.ArrayList; import java.util.ArrayList;
import java.util.Base64; import java.util.Base64;
import java.util.Collections; import java.util.Collections;
@@ -31,14 +29,16 @@ import com.azure.ai.openai.OpenAIClient;
import com.azure.ai.openai.OpenAIClientBuilder; import com.azure.ai.openai.OpenAIClientBuilder;
import com.azure.ai.openai.implementation.accesshelpers.ChatCompletionsOptionsAccessHelper; import com.azure.ai.openai.implementation.accesshelpers.ChatCompletionsOptionsAccessHelper;
import com.azure.ai.openai.models.ChatChoice; import com.azure.ai.openai.models.ChatChoice;
import com.azure.ai.openai.models.ChatCompletionStreamOptions;
import com.azure.ai.openai.models.ChatCompletions; import com.azure.ai.openai.models.ChatCompletions;
import com.azure.ai.openai.models.ChatCompletionsFunctionToolCall; import com.azure.ai.openai.models.ChatCompletionsFunctionToolCall;
import com.azure.ai.openai.models.ChatCompletionsFunctionToolDefinition; import com.azure.ai.openai.models.ChatCompletionsFunctionToolDefinition;
import com.azure.ai.openai.models.ChatCompletionsFunctionToolDefinitionFunction; import com.azure.ai.openai.models.ChatCompletionsFunctionToolDefinitionFunction;
import com.azure.ai.openai.models.ChatCompletionsJsonResponseFormat; import com.azure.ai.openai.models.ChatCompletionsJsonResponseFormat;
import com.azure.ai.openai.models.ChatCompletionsJsonSchemaResponseFormat;
import com.azure.ai.openai.models.ChatCompletionsJsonSchemaResponseFormatJsonSchema;
import com.azure.ai.openai.models.ChatCompletionsOptions; import com.azure.ai.openai.models.ChatCompletionsOptions;
import com.azure.ai.openai.models.ChatCompletionsResponseFormat; import com.azure.ai.openai.models.ChatCompletionsResponseFormat;
import com.azure.ai.openai.models.ChatCompletionStreamOptions;
import com.azure.ai.openai.models.ChatCompletionsTextResponseFormat; import com.azure.ai.openai.models.ChatCompletionsTextResponseFormat;
import com.azure.ai.openai.models.ChatCompletionsToolCall; import com.azure.ai.openai.models.ChatCompletionsToolCall;
import com.azure.ai.openai.models.ChatCompletionsToolDefinition; import com.azure.ai.openai.models.ChatCompletionsToolDefinition;
@@ -55,17 +55,18 @@ import com.azure.ai.openai.models.CompletionsFinishReason;
import com.azure.ai.openai.models.CompletionsUsage; import com.azure.ai.openai.models.CompletionsUsage;
import com.azure.ai.openai.models.ContentFilterResultsForPrompt; import com.azure.ai.openai.models.ContentFilterResultsForPrompt;
import com.azure.ai.openai.models.FunctionCall; import com.azure.ai.openai.models.FunctionCall;
import com.azure.ai.openai.models.ReasoningEffortValue;
import com.azure.core.util.BinaryData; import com.azure.core.util.BinaryData;
import io.micrometer.observation.Observation; import io.micrometer.observation.Observation;
import io.micrometer.observation.ObservationRegistry; import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.contextpropagation.ObservationThreadLocalAccessor; import io.micrometer.observation.contextpropagation.ObservationThreadLocalAccessor;
import org.slf4j.Logger; import org.slf4j.Logger;
import org.slf4j.LoggerFactory; import org.slf4j.LoggerFactory;
import org.springframework.ai.azure.openai.AzureOpenAiResponseFormat.JsonSchema;
import org.springframework.ai.azure.openai.AzureOpenAiResponseFormat.Type;
import reactor.core.publisher.Flux; import reactor.core.publisher.Flux;
import reactor.core.scheduler.Schedulers; import reactor.core.scheduler.Schedulers;
import org.springframework.ai.azure.openai.AzureOpenAiResponseFormat.JsonSchema;
import org.springframework.ai.azure.openai.AzureOpenAiResponseFormat.Type;
import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.ToolResponseMessage; import org.springframework.ai.chat.messages.ToolResponseMessage;
@@ -77,7 +78,6 @@ import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.metadata.PromptMetadata; import org.springframework.ai.chat.metadata.PromptMetadata;
import org.springframework.ai.chat.metadata.PromptMetadata.PromptFilterMetadata; import org.springframework.ai.chat.metadata.PromptMetadata.PromptFilterMetadata;
import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.support.UsageCalculator;
import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation; import org.springframework.ai.chat.model.Generation;
@@ -96,9 +96,11 @@ import org.springframework.ai.model.tool.ToolCallingManager;
import org.springframework.ai.model.tool.ToolExecutionEligibilityPredicate; import org.springframework.ai.model.tool.ToolExecutionEligibilityPredicate;
import org.springframework.ai.model.tool.ToolExecutionResult; import org.springframework.ai.model.tool.ToolExecutionResult;
import org.springframework.ai.observation.conventions.AiProvider; import org.springframework.ai.observation.conventions.AiProvider;
import org.springframework.ai.support.UsageCalculator;
import org.springframework.ai.tool.definition.ToolDefinition; import org.springframework.ai.tool.definition.ToolDefinition;
import org.springframework.util.Assert; import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils; import org.springframework.util.CollectionUtils;
import org.springframework.util.StringUtils;
/** /**
* {@link ChatModel} implementation for {@literal Microsoft Azure AI} backed by * {@link ChatModel} implementation for {@literal Microsoft Azure AI} backed by
@@ -761,6 +763,14 @@ public class AzureOpenAiChatModel implements ChatModel {
mergedAzureOptions.setEnhancements(fromAzureOptions.getEnhancements() != null mergedAzureOptions.setEnhancements(fromAzureOptions.getEnhancements() != null
? fromAzureOptions.getEnhancements() : toSpringAiOptions.getEnhancements()); ? fromAzureOptions.getEnhancements() : toSpringAiOptions.getEnhancements());
ReasoningEffortValue reasoningEffort = (fromAzureOptions.getReasoningEffort() != null)
? fromAzureOptions.getReasoningEffort() : (StringUtils.hasText(toSpringAiOptions.getReasoningEffort())
? ReasoningEffortValue.fromString(toSpringAiOptions.getReasoningEffort()) : null);
if (reasoningEffort != null) {
mergedAzureOptions.setReasoningEffort(reasoningEffort);
}
return mergedAzureOptions; return mergedAzureOptions;
} }
@@ -850,6 +860,11 @@ public class AzureOpenAiChatModel implements ChatModel {
mergedAzureOptions.setEnhancements(fromSpringAiOptions.getEnhancements()); mergedAzureOptions.setEnhancements(fromSpringAiOptions.getEnhancements());
} }
if (StringUtils.hasText(fromSpringAiOptions.getReasoningEffort())) {
mergedAzureOptions
.setReasoningEffort(ReasoningEffortValue.fromString(fromSpringAiOptions.getReasoningEffort()));
}
return mergedAzureOptions; return mergedAzureOptions;
} }
@@ -915,6 +930,10 @@ public class AzureOpenAiChatModel implements ChatModel {
copyOptions.setEnhancements(fromOptions.getEnhancements()); copyOptions.setEnhancements(fromOptions.getEnhancements());
} }
if (fromOptions.getReasoningEffort() != null) {
copyOptions.setReasoningEffort(fromOptions.getReasoningEffort());
}
return copyOptions; return copyOptions;
} }

View File

@@ -207,6 +207,15 @@ public class AzureOpenAiChatOptions implements ToolCallingChatOptions {
@JsonIgnore @JsonIgnore
private Boolean enableStreamUsage; private Boolean enableStreamUsage;
/**
* Constrains effort on reasoning for reasoning models. Currently supported values are
* low, medium, and high. Reducing reasoning effort can result in faster responses and
* fewer tokens used on reasoning in a response. Optional. Defaults to medium. Only
* for reasoning models.
*/
@JsonProperty("reasoning_effort")
private String reasoningEffort;
@Override @Override
@JsonIgnore @JsonIgnore
public List<ToolCallback> getToolCallbacks() { public List<ToolCallback> getToolCallbacks() {
@@ -268,6 +277,7 @@ public class AzureOpenAiChatOptions implements ToolCallingChatOptions {
.toolNames(fromOptions.getToolNames() != null ? new HashSet<>(fromOptions.getToolNames()) : null) .toolNames(fromOptions.getToolNames() != null ? new HashSet<>(fromOptions.getToolNames()) : null)
.responseFormat(fromOptions.getResponseFormat()) .responseFormat(fromOptions.getResponseFormat())
.streamUsage(fromOptions.getStreamUsage()) .streamUsage(fromOptions.getStreamUsage())
.reasoningEffort(fromOptions.getReasoningEffort())
.seed(fromOptions.getSeed()) .seed(fromOptions.getSeed())
.logprobs(fromOptions.isLogprobs()) .logprobs(fromOptions.isLogprobs())
.topLogprobs(fromOptions.getTopLogProbs()) .topLogprobs(fromOptions.getTopLogProbs())
@@ -408,6 +418,14 @@ public class AzureOpenAiChatOptions implements ToolCallingChatOptions {
this.enableStreamUsage = enableStreamUsage; this.enableStreamUsage = enableStreamUsage;
} }
public String getReasoningEffort() {
return this.reasoningEffort;
}
public void setReasoningEffort(String reasoningEffort) {
this.reasoningEffort = reasoningEffort;
}
@Override @Override
@JsonIgnore @JsonIgnore
public Integer getTopK() { public Integer getTopK() {
@@ -490,6 +508,7 @@ public class AzureOpenAiChatOptions implements ToolCallingChatOptions {
&& Objects.equals(this.enhancements, that.enhancements) && Objects.equals(this.enhancements, that.enhancements)
&& Objects.equals(this.streamOptions, that.streamOptions) && Objects.equals(this.streamOptions, that.streamOptions)
&& Objects.equals(this.enableStreamUsage, that.enableStreamUsage) && Objects.equals(this.enableStreamUsage, that.enableStreamUsage)
&& Objects.equals(this.reasoningEffort, that.reasoningEffort)
&& Objects.equals(this.toolContext, that.toolContext) && Objects.equals(this.maxTokens, that.maxTokens) && Objects.equals(this.toolContext, that.toolContext) && Objects.equals(this.maxTokens, that.maxTokens)
&& Objects.equals(this.frequencyPenalty, that.frequencyPenalty) && Objects.equals(this.frequencyPenalty, that.frequencyPenalty)
&& Objects.equals(this.presencePenalty, that.presencePenalty) && Objects.equals(this.presencePenalty, that.presencePenalty)
@@ -500,8 +519,9 @@ public class AzureOpenAiChatOptions implements ToolCallingChatOptions {
public int hashCode() { public int hashCode() {
return Objects.hash(this.logitBias, this.user, this.n, this.stop, this.deploymentName, this.responseFormat, return Objects.hash(this.logitBias, this.user, this.n, this.stop, this.deploymentName, this.responseFormat,
this.toolCallbacks, this.toolNames, this.internalToolExecutionEnabled, this.seed, this.logprobs, this.toolCallbacks, this.toolNames, this.internalToolExecutionEnabled, this.seed, this.logprobs,
this.topLogProbs, this.enhancements, this.streamOptions, this.enableStreamUsage, this.toolContext, this.topLogProbs, this.enhancements, this.streamOptions, this.reasoningEffort, this.enableStreamUsage,
this.maxTokens, this.frequencyPenalty, this.presencePenalty, this.temperature, this.topP); this.toolContext, this.maxTokens, this.frequencyPenalty, this.presencePenalty, this.temperature,
this.topP);
} }
public static class Builder { public static class Builder {
@@ -576,6 +596,11 @@ public class AzureOpenAiChatOptions implements ToolCallingChatOptions {
return this; return this;
} }
public Builder reasoningEffort(String reasoningEffort) {
this.options.reasoningEffort = reasoningEffort;
return this;
}
public Builder seed(Long seed) { public Builder seed(Long seed) {
this.options.seed = seed; this.options.seed = seed;
return this; return this;

View File

@@ -59,6 +59,7 @@ class AzureOpenAiChatOptionsTests {
.user("test-user") .user("test-user")
.responseFormat(responseFormat) .responseFormat(responseFormat)
.streamUsage(true) .streamUsage(true)
.reasoningEffort("low")
.seed(12345L) .seed(12345L)
.logprobs(true) .logprobs(true)
.topLogprobs(5) .topLogprobs(5)
@@ -68,10 +69,10 @@ class AzureOpenAiChatOptionsTests {
assertThat(options) assertThat(options)
.extracting("deploymentName", "frequencyPenalty", "logitBias", "maxTokens", "n", "presencePenalty", "stop", .extracting("deploymentName", "frequencyPenalty", "logitBias", "maxTokens", "n", "presencePenalty", "stop",
"temperature", "topP", "user", "responseFormat", "streamUsage", "seed", "logprobs", "topLogProbs", "temperature", "topP", "user", "responseFormat", "streamUsage", "reasoningEffort", "seed",
"enhancements", "streamOptions") "logprobs", "topLogProbs", "enhancements", "streamOptions")
.containsExactly("test-deployment", 0.5, Map.of("token1", 1, "token2", -1), 200, 2, 0.8, .containsExactly("test-deployment", 0.5, Map.of("token1", 1, "token2", -1), 200, 2, 0.8,
List.of("stop1", "stop2"), 0.7, 0.9, "test-user", responseFormat, true, 12345L, true, 5, List.of("stop1", "stop2"), 0.7, 0.9, "test-user", responseFormat, true, "low", 12345L, true, 5,
enhancements, streamOptions); enhancements, streamOptions);
} }
@@ -100,6 +101,7 @@ class AzureOpenAiChatOptionsTests {
.user("test-user") .user("test-user")
.responseFormat(responseFormat) .responseFormat(responseFormat)
.streamUsage(true) .streamUsage(true)
.reasoningEffort("low")
.seed(12345L) .seed(12345L)
.logprobs(true) .logprobs(true)
.topLogprobs(5) .topLogprobs(5)
@@ -137,6 +139,7 @@ class AzureOpenAiChatOptionsTests {
options.setUser("test-user"); options.setUser("test-user");
options.setResponseFormat(responseFormat); options.setResponseFormat(responseFormat);
options.setStreamUsage(true); options.setStreamUsage(true);
options.setReasoningEffort("low");
options.setSeed(12345L); options.setSeed(12345L);
options.setLogprobs(true); options.setLogprobs(true);
options.setTopLogProbs(5); options.setTopLogProbs(5);
@@ -158,6 +161,7 @@ class AzureOpenAiChatOptionsTests {
assertThat(options.getUser()).isEqualTo("test-user"); assertThat(options.getUser()).isEqualTo("test-user");
assertThat(options.getResponseFormat()).isEqualTo(responseFormat); assertThat(options.getResponseFormat()).isEqualTo(responseFormat);
assertThat(options.getStreamUsage()).isTrue(); assertThat(options.getStreamUsage()).isTrue();
assertThat(options.getReasoningEffort()).isEqualTo("low");
assertThat(options.getSeed()).isEqualTo(12345L); assertThat(options.getSeed()).isEqualTo(12345L);
assertThat(options.isLogprobs()).isTrue(); assertThat(options.isLogprobs()).isTrue();
assertThat(options.getTopLogProbs()).isEqualTo(5); assertThat(options.getTopLogProbs()).isEqualTo(5);
@@ -182,6 +186,7 @@ class AzureOpenAiChatOptionsTests {
assertThat(options.getUser()).isNull(); assertThat(options.getUser()).isNull();
assertThat(options.getResponseFormat()).isNull(); assertThat(options.getResponseFormat()).isNull();
assertThat(options.getStreamUsage()).isNull(); assertThat(options.getStreamUsage()).isNull();
assertThat(options.getReasoningEffort()).isNull();
assertThat(options.getSeed()).isNull(); assertThat(options.getSeed()).isNull();
assertThat(options.isLogprobs()).isNull(); assertThat(options.isLogprobs()).isNull();
assertThat(options.getTopLogProbs()).isNull(); assertThat(options.getTopLogProbs()).isNull();