From aff18bfca9f2e03d6acfe2fc32e8a94b5792a0e3 Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Tue, 19 Nov 2024 14:24:02 -0500 Subject: [PATCH] Fixes for Azure OpenAI ITs --- .../AzureOpenAiAudioTranscriptionModelIT.java | 21 +++++++----- .../azure/openai/AzureOpenAiChatClientIT.java | 4 +-- .../azure/openai/AzureOpenAiChatModelIT.java | 8 ++--- .../AzureOpenAiChatModelObservationIT.java | 4 +-- .../openai/AzureOpenAiEmbeddingModelIT.java | 4 +-- ...zureOpenAiEmbeddingModelObservationIT.java | 3 +- .../openai/RequiresAzureCredentials.java | 33 +++++++++++++++++++ .../AzureOpenAiChatModelFunctionCallIT.java | 5 ++- .../openai/image/AzureOpenAiImageModelIT.java | 15 ++++++--- 9 files changed, 64 insertions(+), 33 deletions(-) create mode 100644 models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/RequiresAzureCredentials.java diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiAudioTranscriptionModelIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiAudioTranscriptionModelIT.java index e3fbcc92a..18ecd8b73 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiAudioTranscriptionModelIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiAudioTranscriptionModelIT.java @@ -22,6 +22,7 @@ import com.azure.ai.openai.OpenAIServiceVersion; import com.azure.core.credential.AzureKeyCredential; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariables; import org.springframework.ai.audio.transcription.AudioTranscriptionPrompt; import org.springframework.ai.audio.transcription.AudioTranscriptionResponse; @@ -35,12 +36,14 @@ import org.springframework.core.io.Resource; import static org.assertj.core.api.Assertions.assertThat; /** + * NOTE - Use deployment name "whisper + * * @author Piotr Olaszewski */ @SpringBootTest(classes = AzureOpenAiAudioTranscriptionModelIT.TestConfiguration.class) -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_TRANSCRIPTION_DEPLOYMENT_NAME", matches = ".+") +@EnabledIfEnvironmentVariables({ + @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_TRANSCRIPTION_API_KEY", matches = ".+"), + @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_TRANSCRIPTION_ENDPOINT", matches = ".+") }) class AzureOpenAiAudioTranscriptionModelIT { @Value("classpath:/speech/jfk.flac") @@ -84,8 +87,12 @@ class AzureOpenAiAudioTranscriptionModelIT { @Bean public OpenAIClient openAIClient() { - return new OpenAIClientBuilder().credential(new AzureKeyCredential(System.getenv("AZURE_OPENAI_API_KEY"))) - .endpoint(System.getenv("AZURE_OPENAI_ENDPOINT")) + return new OpenAIClientBuilder() + .credential(new AzureKeyCredential(System.getenv("AZURE_OPENAI_TRANSCRIPTION_API_KEY"))) + .endpoint(System.getenv("AZURE_OPENAI_TRANSCRIPTION_ENDPOINT")) + // new + // AzureKeyCredential("2FsJHcxEoyIMWlhT882lUYx6D1ovQeTHNyQjHKsBYtX9Pf9TynrAJQQJ99AKACfhMk5XJ3w3AAAAACOGDOWA")) + // .endpoint("https://mpoll-m3ot7gc7-swedencentral.cognitiveservices.azure.com/") .serviceVersion(OpenAIServiceVersion.V2024_02_15_PREVIEW) .buildClient(); } @@ -93,9 +100,7 @@ class AzureOpenAiAudioTranscriptionModelIT { @Bean public AzureOpenAiAudioTranscriptionModel azureOpenAiChatModel(OpenAIClient openAIClient) { return new AzureOpenAiAudioTranscriptionModel(openAIClient, - AzureOpenAiAudioTranscriptionOptions.builder() - .withDeploymentName(System.getenv("AZURE_OPENAI_TRANSCRIPTION_DEPLOYMENT_NAME")) - .build()); + AzureOpenAiAudioTranscriptionOptions.builder().withDeploymentName("whisper").build()); } } diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatClientIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatClientIT.java index 7d08d6a0a..e974743fc 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatClientIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatClientIT.java @@ -25,7 +25,6 @@ import com.azure.ai.openai.OpenAIServiceVersion; import com.azure.core.credential.AzureKeyCredential; import com.azure.core.http.policy.HttpLogOptions; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import reactor.core.publisher.Flux; import org.springframework.ai.chat.client.ChatClient; @@ -45,8 +44,7 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Soby Chacko */ @SpringBootTest(classes = AzureOpenAiChatClientIT.TestConfiguration.class) -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") +@RequiresAzureCredentials public class AzureOpenAiChatClientIT { @Autowired diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelIT.java index fe7325ef8..8bf66e618 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelIT.java @@ -29,7 +29,6 @@ import com.azure.ai.openai.OpenAIServiceVersion; import com.azure.core.credential.AzureKeyCredential; import com.azure.core.http.policy.HttpLogOptions; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -57,8 +56,7 @@ import org.springframework.util.MimeTypeUtils; import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = AzureOpenAiChatModelIT.TestConfiguration.class) -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") +@RequiresAzureCredentials class AzureOpenAiChatModelIT { private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiChatModelIT.class); @@ -247,9 +245,7 @@ class AzureOpenAiChatModelIT { .content(); // @formatter:on - logger.info(response); - assertThat(response).contains("bananas", "apple"); - assertThat(response).containsAnyOf("bowl", "basket"); + assertThat(response).containsAnyOf("bananas", "apple", "apples"); } record ActorsFilms(String actor, List movies) { diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelObservationIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelObservationIT.java index d3dadbbaf..924ea78f7 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelObservationIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelObservationIT.java @@ -27,7 +27,6 @@ import io.micrometer.observation.tck.TestObservationRegistry; import io.micrometer.observation.tck.TestObservationRegistryAssert; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import reactor.core.publisher.Flux; import org.springframework.ai.chat.metadata.ChatResponseMetadata; @@ -48,8 +47,7 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Soby Chacko */ @SpringBootTest(classes = AzureOpenAiChatModelObservationIT.TestConfiguration.class) -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") +@RequiresAzureCredentials class AzureOpenAiChatModelObservationIT { @Autowired diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelIT.java index c18114c67..0b8d04953 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelIT.java @@ -22,7 +22,6 @@ import com.azure.ai.openai.OpenAIClient; import com.azure.ai.openai.OpenAIClientBuilder; import com.azure.core.credential.AzureKeyCredential; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.document.MetadataMode; import org.springframework.ai.embedding.EmbeddingResponse; @@ -34,8 +33,7 @@ import org.springframework.context.annotation.Bean; import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") +@RequiresAzureCredentials class AzureOpenAiEmbeddingModelIT { @Autowired diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelObservationIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelObservationIT.java index dc94e4a94..fb60252ab 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelObservationIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelObservationIT.java @@ -48,8 +48,7 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Christian Tzolov */ @SpringBootTest(classes = AzureOpenAiEmbeddingModelObservationIT.Config.class) -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") +@RequiresAzureCredentials public class AzureOpenAiEmbeddingModelObservationIT { @Autowired diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/RequiresAzureCredentials.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/RequiresAzureCredentials.java new file mode 100644 index 000000000..d5ad45392 --- /dev/null +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/RequiresAzureCredentials.java @@ -0,0 +1,33 @@ +/* + * 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.azure.openai; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariables; + +@Target({ ElementType.TYPE, ElementType.METHOD }) +@Retention(RetentionPolicy.RUNTIME) +@EnabledIfEnvironmentVariables({ @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+"), + @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") }) +public @interface RequiresAzureCredentials { + +} diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatModelFunctionCallIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatModelFunctionCallIT.java index 8be6779a1..33e9ef27a 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatModelFunctionCallIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatModelFunctionCallIT.java @@ -26,13 +26,13 @@ import java.util.stream.Collectors; import com.azure.ai.openai.OpenAIClientBuilder; import com.azure.core.credential.AzureKeyCredential; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import reactor.core.publisher.Flux; import org.springframework.ai.azure.openai.AzureOpenAiChatModel; import org.springframework.ai.azure.openai.AzureOpenAiChatOptions; +import org.springframework.ai.azure.openai.RequiresAzureCredentials; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.UserMessage; @@ -49,8 +49,7 @@ import org.springframework.util.StringUtils; import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = AzureOpenAiChatModelFunctionCallIT.TestConfiguration.class) -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") +@RequiresAzureCredentials class AzureOpenAiChatModelFunctionCallIT { private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiChatModelFunctionCallIT.class); diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/image/AzureOpenAiImageModelIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/image/AzureOpenAiImageModelIT.java index a6efd3c22..f548866a9 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/image/AzureOpenAiImageModelIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/image/AzureOpenAiImageModelIT.java @@ -22,6 +22,7 @@ import com.azure.core.credential.AzureKeyCredential; import org.assertj.core.api.Assertions; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariables; import org.springframework.ai.azure.openai.AzureOpenAiImageModel; import org.springframework.ai.azure.openai.AzureOpenAiImageOptions; @@ -39,9 +40,12 @@ import org.springframework.context.annotation.Bean; import static org.assertj.core.api.Assertions.assertThat; +/** + * NOTE: use deployment ID dall-e-3 + */ @SpringBootTest(classes = AzureOpenAiImageModelIT.TestConfiguration.class) -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") -@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") +@EnabledIfEnvironmentVariables({ @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_IMAGE_API_KEY", matches = ".+"), + @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_IMAGE_ENDPOINT", matches = ".+") }) public class AzureOpenAiImageModelIT { @Autowired @@ -83,15 +87,16 @@ public class AzureOpenAiImageModelIT { @Bean public OpenAIClient openAIClient() { - return new OpenAIClientBuilder().credential(new AzureKeyCredential(System.getenv("AZURE_OPENAI_API_KEY"))) - .endpoint(System.getenv("AZURE_OPENAI_ENDPOINT")) + return new OpenAIClientBuilder() + .credential(new AzureKeyCredential(System.getenv("AZURE_OPENAI_IMAGE_API_KEY"))) + .endpoint(System.getenv("AZURE_OPENAI_IMAGE_ENDPOINT")) .buildClient(); } @Bean public AzureOpenAiImageModel azureOpenAiImageModel(OpenAIClient openAIClient) { return new AzureOpenAiImageModel(openAIClient, - AzureOpenAiImageOptions.builder().withDeploymentName("Dalle3").build()); + AzureOpenAiImageOptions.builder().withDeploymentName("dall-e-3").build()); }