From fefc6f4173ed33240701c915efa31964dc64f603 Mon Sep 17 00:00:00 2001 From: Hyune-c Date: Mon, 1 Apr 2024 21:53:52 +0900 Subject: [PATCH] fix: AzureOpenAiChatOptions to handle null values for presencePenalty and frequencyPenalty --- .../azure/openai/AzureOpenAiChatOptions.java | 8 +++- .../AzureChatCompletionsOptionsTests.java | 40 +++++++++++++++++++ 2 files changed, 46 insertions(+), 2 deletions(-) diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java index fdc6a2354..31df9633f 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java @@ -174,7 +174,9 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio } public Builder withFrequencyPenalty(Float frequencyPenalty) { - this.options.frequencyPenalty = frequencyPenalty.doubleValue(); + if(frequencyPenalty != null) { + this.options.frequencyPenalty = frequencyPenalty.doubleValue(); + } return this; } @@ -194,7 +196,9 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio } public Builder withPresencePenalty(Float presencePenalty) { - this.options.presencePenalty = presencePenalty.doubleValue(); + if(presencePenalty != null) { + this.options.presencePenalty = presencePenalty.doubleValue(); + } return this; } diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureChatCompletionsOptionsTests.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureChatCompletionsOptionsTests.java index 70a0ac67e..fa0726882 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureChatCompletionsOptionsTests.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureChatCompletionsOptionsTests.java @@ -17,10 +17,15 @@ package org.springframework.ai.azure.openai; import com.azure.ai.openai.OpenAIClient; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; import org.mockito.Mockito; import org.springframework.ai.chat.prompt.Prompt; +import java.util.stream.Stream; + import static org.assertj.core.api.Assertions.assertThat; /** @@ -51,4 +56,39 @@ public class AzureChatCompletionsOptionsTests { assertThat(requestOptions.getTemperature()).isEqualTo(99.9f); } + private static Stream providePresencePenaltyAndFrequencyPenaltyTest() { + return Stream.of( + Arguments.of(0.0f, 0.0f), + Arguments.of(0.0f, 1.0f), + Arguments.of(1.0f, 0.0f), + Arguments.of(1.0f, 1.0f), + Arguments.of(1.0f, null), + Arguments.of(null, 1.0f), + Arguments.of(null, null) + ); + } + + @ParameterizedTest + @MethodSource("providePresencePenaltyAndFrequencyPenaltyTest") + public void createChatOptionsWithPresencePenaltyAndFrequencyPenalty(Float presencePenalty, Float frequencyPenalty) { + var options = AzureOpenAiChatOptions.builder() + .withMaxTokens(800) + .withTemperature(0.7F) + .withTopP(0.95F) + .withPresencePenalty(presencePenalty) + .withFrequencyPenalty(frequencyPenalty) + .build(); + + if (presencePenalty == null) { + assertThat(options.getPresencePenalty()).isEqualTo(null); + } else { + assertThat(options.getPresencePenalty().floatValue()).isEqualTo(presencePenalty); + } + + if (frequencyPenalty == null) { + assertThat(options.getFrequencyPenalty()).isEqualTo(null); + } else { + assertThat(options.getFrequencyPenalty().floatValue()).isEqualTo(frequencyPenalty); + } + } }