Improve Azure OpenAI options merging logic
This commit is contained in:
@@ -277,120 +277,135 @@ public class AzureOpenAiChatClient
|
||||
return list != null ? list : Collections.emptyList();
|
||||
}
|
||||
|
||||
// JSON merge doesn't due to Azure OpenAI service bug:
|
||||
// https://github.com/Azure/azure-sdk-for-java/issues/38183
|
||||
private ChatCompletionsOptions merge(ChatCompletionsOptions azureOptions, AzureOpenAiChatOptions springAiOptions) {
|
||||
/**
|
||||
* Merges the Azure's {@link ChatCompletionsOptions} (fromAzureOptions) into the
|
||||
* Spring AI's {@link AzureOpenAiChatOptions} (toSpringAiOptions) and return a new
|
||||
* {@link ChatCompletionsOptions} instance.
|
||||
*/
|
||||
private ChatCompletionsOptions merge(ChatCompletionsOptions fromAzureOptions,
|
||||
AzureOpenAiChatOptions toSpringAiOptions) {
|
||||
|
||||
if (springAiOptions == null) {
|
||||
return azureOptions;
|
||||
if (toSpringAiOptions == null) {
|
||||
return fromAzureOptions;
|
||||
}
|
||||
|
||||
ChatCompletionsOptions mergedAzureOptions = new ChatCompletionsOptions(azureOptions.getMessages());
|
||||
mergedAzureOptions.setStream(azureOptions.isStream());
|
||||
ChatCompletionsOptions mergedAzureOptions = new ChatCompletionsOptions(fromAzureOptions.getMessages());
|
||||
mergedAzureOptions.setStream(fromAzureOptions.isStream());
|
||||
|
||||
mergedAzureOptions.setMaxTokens(
|
||||
(azureOptions.getMaxTokens() != null) ? azureOptions.getMaxTokens() : springAiOptions.getMaxTokens());
|
||||
mergedAzureOptions.setMaxTokens((fromAzureOptions.getMaxTokens() != null) ? fromAzureOptions.getMaxTokens()
|
||||
: toSpringAiOptions.getMaxTokens());
|
||||
|
||||
mergedAzureOptions.setLogitBias(
|
||||
azureOptions.getLogitBias() != null ? azureOptions.getLogitBias() : springAiOptions.getLogitBias());
|
||||
|
||||
mergedAzureOptions.setStop(azureOptions.getStop() != null ? azureOptions.getStop() : springAiOptions.getStop());
|
||||
|
||||
mergedAzureOptions.setTemperature(azureOptions.getTemperature());
|
||||
if (mergedAzureOptions.getTemperature() == null && springAiOptions.getTemperature() != null) {
|
||||
mergedAzureOptions.setTemperature(springAiOptions.getTemperature().doubleValue());
|
||||
}
|
||||
|
||||
mergedAzureOptions.setTopP(azureOptions.getTopP());
|
||||
if (mergedAzureOptions.getTopP() == null && springAiOptions.getTopP() != null) {
|
||||
mergedAzureOptions.setTopP(springAiOptions.getTopP().doubleValue());
|
||||
}
|
||||
|
||||
mergedAzureOptions.setFrequencyPenalty(azureOptions.getFrequencyPenalty());
|
||||
if (mergedAzureOptions.getFrequencyPenalty() == null && springAiOptions.getFrequencyPenalty() != null) {
|
||||
mergedAzureOptions.setFrequencyPenalty(springAiOptions.getFrequencyPenalty().doubleValue());
|
||||
}
|
||||
|
||||
mergedAzureOptions.setPresencePenalty(azureOptions.getPresencePenalty());
|
||||
if (mergedAzureOptions.getPresencePenalty() == null && springAiOptions.getPresencePenalty() != null) {
|
||||
mergedAzureOptions.setPresencePenalty(springAiOptions.getPresencePenalty().doubleValue());
|
||||
}
|
||||
|
||||
mergedAzureOptions.setN(azureOptions.getN() != null ? azureOptions.getN() : springAiOptions.getN());
|
||||
|
||||
mergedAzureOptions.setUser(azureOptions.getUser() != null ? azureOptions.getUser() : springAiOptions.getUser());
|
||||
mergedAzureOptions.setLogitBias(fromAzureOptions.getLogitBias() != null ? fromAzureOptions.getLogitBias()
|
||||
: toSpringAiOptions.getLogitBias());
|
||||
|
||||
mergedAzureOptions
|
||||
.setModel(azureOptions.getModel() != null ? azureOptions.getModel() : springAiOptions.getDeploymentName());
|
||||
.setStop(fromAzureOptions.getStop() != null ? fromAzureOptions.getStop() : toSpringAiOptions.getStop());
|
||||
|
||||
mergedAzureOptions.setTemperature(fromAzureOptions.getTemperature());
|
||||
if (mergedAzureOptions.getTemperature() == null && toSpringAiOptions.getTemperature() != null) {
|
||||
mergedAzureOptions.setTemperature(toSpringAiOptions.getTemperature().doubleValue());
|
||||
}
|
||||
|
||||
mergedAzureOptions.setTopP(fromAzureOptions.getTopP());
|
||||
if (mergedAzureOptions.getTopP() == null && toSpringAiOptions.getTopP() != null) {
|
||||
mergedAzureOptions.setTopP(toSpringAiOptions.getTopP().doubleValue());
|
||||
}
|
||||
|
||||
mergedAzureOptions.setFrequencyPenalty(fromAzureOptions.getFrequencyPenalty());
|
||||
if (mergedAzureOptions.getFrequencyPenalty() == null && toSpringAiOptions.getFrequencyPenalty() != null) {
|
||||
mergedAzureOptions.setFrequencyPenalty(toSpringAiOptions.getFrequencyPenalty().doubleValue());
|
||||
}
|
||||
|
||||
mergedAzureOptions.setPresencePenalty(fromAzureOptions.getPresencePenalty());
|
||||
if (mergedAzureOptions.getPresencePenalty() == null && toSpringAiOptions.getPresencePenalty() != null) {
|
||||
mergedAzureOptions.setPresencePenalty(toSpringAiOptions.getPresencePenalty().doubleValue());
|
||||
}
|
||||
|
||||
mergedAzureOptions.setN(fromAzureOptions.getN() != null ? fromAzureOptions.getN() : toSpringAiOptions.getN());
|
||||
|
||||
mergedAzureOptions
|
||||
.setUser(fromAzureOptions.getUser() != null ? fromAzureOptions.getUser() : toSpringAiOptions.getUser());
|
||||
|
||||
mergedAzureOptions.setModel(fromAzureOptions.getModel() != null ? fromAzureOptions.getModel()
|
||||
: toSpringAiOptions.getDeploymentName());
|
||||
|
||||
return mergedAzureOptions;
|
||||
}
|
||||
|
||||
// JSON merge doesn't due to Azure OpenAI service bug:
|
||||
// https://github.com/Azure/azure-sdk-for-java/issues/38183
|
||||
private ChatCompletionsOptions merge(AzureOpenAiChatOptions springAiOptions, ChatCompletionsOptions azureOptions) {
|
||||
if (springAiOptions == null) {
|
||||
return azureOptions;
|
||||
/**
|
||||
* Merges the {@link AzureOpenAiChatOptions}, fromSpringAiOptions, into the
|
||||
* {@link ChatCompletionsOptions}, toAzureOptions, and returns a new
|
||||
* {@link ChatCompletionsOptions} instance.
|
||||
* @param fromSpringAiOptions the {@link AzureOpenAiChatOptions} to merge from.
|
||||
* @param toAzureOptions the {@link ChatCompletionsOptions} to merge to.
|
||||
* @return a new {@link ChatCompletionsOptions} instance.
|
||||
*/
|
||||
private ChatCompletionsOptions merge(AzureOpenAiChatOptions fromSpringAiOptions,
|
||||
ChatCompletionsOptions toAzureOptions) {
|
||||
|
||||
if (fromSpringAiOptions == null) {
|
||||
return toAzureOptions;
|
||||
}
|
||||
|
||||
ChatCompletionsOptions mergedAzureOptions = new ChatCompletionsOptions(azureOptions.getMessages());
|
||||
mergedAzureOptions = merge(azureOptions, mergedAzureOptions);
|
||||
ChatCompletionsOptions mergedAzureOptions = this.copy(toAzureOptions);
|
||||
|
||||
mergedAzureOptions.setStream(azureOptions.isStream());
|
||||
|
||||
if (springAiOptions.getMaxTokens() != null) {
|
||||
mergedAzureOptions.setMaxTokens(springAiOptions.getMaxTokens());
|
||||
if (fromSpringAiOptions.getMaxTokens() != null) {
|
||||
mergedAzureOptions.setMaxTokens(fromSpringAiOptions.getMaxTokens());
|
||||
}
|
||||
|
||||
if (springAiOptions.getLogitBias() != null) {
|
||||
mergedAzureOptions.setLogitBias(springAiOptions.getLogitBias());
|
||||
if (fromSpringAiOptions.getLogitBias() != null) {
|
||||
mergedAzureOptions.setLogitBias(fromSpringAiOptions.getLogitBias());
|
||||
}
|
||||
|
||||
if (springAiOptions.getStop() != null) {
|
||||
mergedAzureOptions.setStop(springAiOptions.getStop());
|
||||
if (fromSpringAiOptions.getStop() != null) {
|
||||
mergedAzureOptions.setStop(fromSpringAiOptions.getStop());
|
||||
}
|
||||
|
||||
if (springAiOptions.getTemperature() != null && springAiOptions.getTemperature() != null) {
|
||||
mergedAzureOptions.setTemperature(springAiOptions.getTemperature().doubleValue());
|
||||
if (fromSpringAiOptions.getTemperature() != null) {
|
||||
mergedAzureOptions.setTemperature(fromSpringAiOptions.getTemperature().doubleValue());
|
||||
}
|
||||
|
||||
if (springAiOptions.getTopP() != null && springAiOptions.getTopP() != null) {
|
||||
mergedAzureOptions.setTopP(springAiOptions.getTopP().doubleValue());
|
||||
if (fromSpringAiOptions.getTopP() != null) {
|
||||
mergedAzureOptions.setTopP(fromSpringAiOptions.getTopP().doubleValue());
|
||||
}
|
||||
|
||||
if (springAiOptions.getFrequencyPenalty() != null && springAiOptions.getFrequencyPenalty() != null) {
|
||||
mergedAzureOptions.setFrequencyPenalty(springAiOptions.getFrequencyPenalty().doubleValue());
|
||||
if (fromSpringAiOptions.getFrequencyPenalty() != null) {
|
||||
mergedAzureOptions.setFrequencyPenalty(fromSpringAiOptions.getFrequencyPenalty().doubleValue());
|
||||
}
|
||||
|
||||
if (springAiOptions.getPresencePenalty() != null && springAiOptions.getPresencePenalty() != null) {
|
||||
mergedAzureOptions.setPresencePenalty(springAiOptions.getPresencePenalty().doubleValue());
|
||||
if (fromSpringAiOptions.getPresencePenalty() != null) {
|
||||
mergedAzureOptions.setPresencePenalty(fromSpringAiOptions.getPresencePenalty().doubleValue());
|
||||
}
|
||||
|
||||
if (springAiOptions.getN() != null) {
|
||||
mergedAzureOptions.setN(springAiOptions.getN());
|
||||
if (fromSpringAiOptions.getN() != null) {
|
||||
mergedAzureOptions.setN(fromSpringAiOptions.getN());
|
||||
}
|
||||
|
||||
if (springAiOptions.getUser() != null) {
|
||||
mergedAzureOptions.setUser(springAiOptions.getUser());
|
||||
if (fromSpringAiOptions.getUser() != null) {
|
||||
mergedAzureOptions.setUser(fromSpringAiOptions.getUser());
|
||||
}
|
||||
|
||||
if (springAiOptions.getDeploymentName() != null) {
|
||||
mergedAzureOptions.setModel(springAiOptions.getDeploymentName());
|
||||
if (fromSpringAiOptions.getDeploymentName() != null) {
|
||||
mergedAzureOptions.setModel(fromSpringAiOptions.getDeploymentName());
|
||||
}
|
||||
|
||||
return mergedAzureOptions;
|
||||
}
|
||||
|
||||
// https://github.com/Azure/azure-sdk-for-java/blob/azure-ai-openai_1.0.0-beta.6/sdk/openai/azure-ai-openai/src/samples/java/com/azure/ai/openai/usage/GetChatCompletionsToolCallSample.java
|
||||
|
||||
/**
|
||||
* Merges the fromOptions into the toOptions and returns a new ChatCompletionsOptions
|
||||
* instance.
|
||||
* @param fromOptions the ChatCompletionsOptions to merge from.
|
||||
* @param toOptions the ChatCompletionsOptions to merge to.
|
||||
* @return a new ChatCompletionsOptions instance.
|
||||
*/
|
||||
private ChatCompletionsOptions merge(ChatCompletionsOptions fromOptions, ChatCompletionsOptions toOptions) {
|
||||
|
||||
if (fromOptions == null) {
|
||||
return toOptions;
|
||||
}
|
||||
|
||||
ChatCompletionsOptions mergedOptions = new ChatCompletionsOptions(toOptions.getMessages());
|
||||
mergedOptions.setStream(toOptions.isStream());
|
||||
ChatCompletionsOptions mergedOptions = this.copy(toOptions);
|
||||
|
||||
if (fromOptions.getMaxTokens() != null) {
|
||||
mergedOptions.setMaxTokens(fromOptions.getMaxTokens());
|
||||
@@ -426,6 +441,50 @@ public class AzureOpenAiChatClient
|
||||
return mergedOptions;
|
||||
}
|
||||
|
||||
/**
|
||||
* Copy the fromOptions into a new ChatCompletionsOptions instance.
|
||||
* @param fromOptions the ChatCompletionsOptions to copy from.
|
||||
* @return a new ChatCompletionsOptions instance.
|
||||
*/
|
||||
private ChatCompletionsOptions copy(ChatCompletionsOptions fromOptions) {
|
||||
|
||||
ChatCompletionsOptions copyOptions = new ChatCompletionsOptions(fromOptions.getMessages());
|
||||
copyOptions.setStream(fromOptions.isStream());
|
||||
|
||||
if (fromOptions.getMaxTokens() != null) {
|
||||
copyOptions.setMaxTokens(fromOptions.getMaxTokens());
|
||||
}
|
||||
if (fromOptions.getLogitBias() != null) {
|
||||
copyOptions.setLogitBias(fromOptions.getLogitBias());
|
||||
}
|
||||
if (fromOptions.getStop() != null) {
|
||||
copyOptions.setStop(fromOptions.getStop());
|
||||
}
|
||||
if (fromOptions.getTemperature() != null) {
|
||||
copyOptions.setTemperature(fromOptions.getTemperature());
|
||||
}
|
||||
if (fromOptions.getTopP() != null) {
|
||||
copyOptions.setTopP(fromOptions.getTopP());
|
||||
}
|
||||
if (fromOptions.getFrequencyPenalty() != null) {
|
||||
copyOptions.setFrequencyPenalty(fromOptions.getFrequencyPenalty());
|
||||
}
|
||||
if (fromOptions.getPresencePenalty() != null) {
|
||||
copyOptions.setPresencePenalty(fromOptions.getPresencePenalty());
|
||||
}
|
||||
if (fromOptions.getN() != null) {
|
||||
copyOptions.setN(fromOptions.getN());
|
||||
}
|
||||
if (fromOptions.getUser() != null) {
|
||||
copyOptions.setUser(fromOptions.getUser());
|
||||
}
|
||||
if (fromOptions.getModel() != null) {
|
||||
copyOptions.setModel(fromOptions.getModel());
|
||||
}
|
||||
|
||||
return copyOptions;
|
||||
}
|
||||
|
||||
@Override
|
||||
protected ChatCompletionsOptions doCreateToolResponseRequest(ChatCompletionsOptions previousRequest,
|
||||
ChatRequestMessage responseMessage, List<ChatRequestMessage> conversationHistory) {
|
||||
|
||||
@@ -24,6 +24,8 @@ import org.mockito.Mockito;
|
||||
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.stream.Stream;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
@@ -37,8 +39,21 @@ public class AzureChatCompletionsOptionsTests {
|
||||
public void createRequestWithChatOptions() {
|
||||
|
||||
OpenAIClient mockClient = Mockito.mock(OpenAIClient.class);
|
||||
var client = new AzureOpenAiChatClient(mockClient,
|
||||
AzureOpenAiChatOptions.builder().withDeploymentName("DEFAULT_MODEL").withTemperature(66.6f).build());
|
||||
|
||||
var defaultOptions = AzureOpenAiChatOptions.builder()
|
||||
.withDeploymentName("DEFAULT_MODEL")
|
||||
.withTemperature(66.6f)
|
||||
.withFrequencyPenalty(696.9f)
|
||||
.withPresencePenalty(969.6f)
|
||||
.withLogitBias(Map.of("foo", 1))
|
||||
.withMaxTokens(969)
|
||||
.withN(69)
|
||||
.withStop(List.of("foo", "bar"))
|
||||
.withTopP(0.69f)
|
||||
.withUser("user")
|
||||
.build();
|
||||
|
||||
var client = new AzureOpenAiChatClient(mockClient, defaultOptions);
|
||||
|
||||
var requestOptions = client.toAzureChatCompletionsOptions(new Prompt("Test message content"));
|
||||
|
||||
@@ -46,14 +61,42 @@ public class AzureChatCompletionsOptionsTests {
|
||||
|
||||
assertThat(requestOptions.getModel()).isEqualTo("DEFAULT_MODEL");
|
||||
assertThat(requestOptions.getTemperature()).isEqualTo(66.6f);
|
||||
assertThat(requestOptions.getFrequencyPenalty()).isEqualTo(696.9f);
|
||||
assertThat(requestOptions.getPresencePenalty()).isEqualTo(969.6f);
|
||||
assertThat(requestOptions.getLogitBias()).isEqualTo(Map.of("foo", 1));
|
||||
assertThat(requestOptions.getMaxTokens()).isEqualTo(969);
|
||||
assertThat(requestOptions.getN()).isEqualTo(69);
|
||||
assertThat(requestOptions.getStop()).isEqualTo(List.of("foo", "bar"));
|
||||
assertThat(requestOptions.getTopP()).isEqualTo(0.69f);
|
||||
assertThat(requestOptions.getUser()).isEqualTo("user");
|
||||
|
||||
requestOptions = client.toAzureChatCompletionsOptions(new Prompt("Test message content",
|
||||
AzureOpenAiChatOptions.builder().withDeploymentName("PROMPT_MODEL").withTemperature(99.9f).build()));
|
||||
var runtimeOptions = AzureOpenAiChatOptions.builder()
|
||||
.withDeploymentName("PROMPT_MODEL")
|
||||
.withTemperature(99.9f)
|
||||
.withFrequencyPenalty(100f)
|
||||
.withPresencePenalty(100f)
|
||||
.withLogitBias(Map.of("foo", 2))
|
||||
.withMaxTokens(100)
|
||||
.withN(100)
|
||||
.withStop(List.of("foo", "bar"))
|
||||
.withTopP(0.111f)
|
||||
.withUser("user2")
|
||||
.build();
|
||||
|
||||
requestOptions = client.toAzureChatCompletionsOptions(new Prompt("Test message content", runtimeOptions));
|
||||
|
||||
assertThat(requestOptions.getMessages()).hasSize(1);
|
||||
|
||||
assertThat(requestOptions.getModel()).isEqualTo("PROMPT_MODEL");
|
||||
assertThat(requestOptions.getTemperature()).isEqualTo(99.9f);
|
||||
assertThat(requestOptions.getFrequencyPenalty()).isEqualTo(100f);
|
||||
assertThat(requestOptions.getPresencePenalty()).isEqualTo(100f);
|
||||
assertThat(requestOptions.getLogitBias()).isEqualTo(Map.of("foo", 2));
|
||||
assertThat(requestOptions.getMaxTokens()).isEqualTo(100);
|
||||
assertThat(requestOptions.getN()).isEqualTo(100);
|
||||
assertThat(requestOptions.getStop()).isEqualTo(List.of("foo", "bar"));
|
||||
assertThat(requestOptions.getTopP()).isEqualTo(0.111f);
|
||||
assertThat(requestOptions.getUser()).isEqualTo("user2");
|
||||
}
|
||||
|
||||
private static Stream<Arguments> providePresencePenaltyAndFrequencyPenaltyTest() {
|
||||
|
||||
Reference in New Issue
Block a user