From 14e4a6cccf6027ab1499f2f6b6240096213c8808 Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Wed, 21 Feb 2024 13:47:11 -0500 Subject: [PATCH] StabilityAI Autoconfiguration improvements * Add IT Test * Add starter to Spring AI bom --- pom.xml | 1 + spring-ai-bom/pom.xml | 8 +- .../StabilityAiImageAutoConfiguration.java | 25 ++++-- .../StabilityAiImageProperties.java | 6 +- ...ockCohereEmbeddingAutoConfigurationIT.java | 9 +-- .../StabilityAiAutoConfigurationIT.java | 42 ++++++++++ .../StabilityAiImagePropertiesTests.java | 77 +++++++++++++++++++ 7 files changed, 153 insertions(+), 15 deletions(-) create mode 100644 spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiAutoConfigurationIT.java create mode 100644 spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImagePropertiesTests.java diff --git a/pom.xml b/pom.xml index cf06a4338..436999542 100644 --- a/pom.xml +++ b/pom.xml @@ -37,6 +37,7 @@ spring-ai-spring-boot-starters/spring-ai-starter-azure-store spring-ai-spring-boot-starters/spring-ai-starter-weaviate-store spring-ai-spring-boot-starters/spring-ai-starter-redis + spring-ai-spring-boot-starters/spring-ai-starter-stability-ai spring-ai-spring-boot-starters/spring-ai-starter-neo4j-store spring-ai-spring-boot-starters/spring-ai-starter-postgresml-embedding spring-ai-docs diff --git a/spring-ai-bom/pom.xml b/spring-ai-bom/pom.xml index 8ad3c6a3d..dbcc33bc1 100644 --- a/spring-ai-bom/pom.xml +++ b/spring-ai-bom/pom.xml @@ -163,7 +163,7 @@ ${project.version} - + org.springframework.ai spring-ai-azure-openai-spring-boot-starter @@ -236,6 +236,12 @@ ${project.version} + + org.springframework.ai + spring-ai-stability-ai-spring-boot-starter + ${project.version} + + org.springframework.ai spring-ai-transformers-spring-boot-starter diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageAutoConfiguration.java index 60f0b3468..24a3377cc 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageAutoConfiguration.java @@ -17,27 +17,40 @@ package org.springframework.ai.autoconfigure.stabilityai; import org.springframework.ai.stabilityai.StabilityAiImageClient; import org.springframework.ai.stabilityai.api.StabilityAiApi; +import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; +import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; import org.springframework.boot.context.properties.EnableConfigurationProperties; import org.springframework.context.annotation.Bean; -import org.springframework.web.client.RestClient; +import org.springframework.util.Assert; +import org.springframework.util.StringUtils; /** * @author Mark Pollack * @since 0.8.0 */ +@AutoConfiguration(after = { RestClientAutoConfiguration.class }) @ConditionalOnClass(StabilityAiApi.class) -@EnableConfigurationProperties({ StabilityAiImageProperties.class }) +@EnableConfigurationProperties({ StabilityAiConnectionProperties.class, StabilityAiImageProperties.class }) public class StabilityAiImageAutoConfiguration { @Bean @ConditionalOnMissingBean - public StabilityAiApi stabilityAiApi(StabilityAiImageProperties stabilityAiImageProperties, - RestClient.Builder restClientBuilder) { - return new StabilityAiApi(stabilityAiImageProperties.getApiKey(), stabilityAiImageProperties.getBaseUrl(), - stabilityAiImageProperties.getOptions().getModel(), restClientBuilder); + public StabilityAiApi stabilityAiApi(StabilityAiConnectionProperties commonProperties, + StabilityAiImageProperties imageProperties) { + + String apiKey = StringUtils.hasText(imageProperties.getApiKey()) ? imageProperties.getApiKey() + : commonProperties.getApiKey(); + + String baseUrl = StringUtils.hasText(imageProperties.getBaseUrl()) ? imageProperties.getBaseUrl() + : commonProperties.getBaseUrl(); + + Assert.hasText(apiKey, "StabilityAI API key must be set"); + Assert.hasText(baseUrl, "StabilityAI base URL must be set"); + + return new StabilityAiApi(apiKey, imageProperties.getOptions().getModel(), baseUrl); } @Bean diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageProperties.java index 9b3286cf0..ce322a591 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageProperties.java @@ -15,7 +15,6 @@ */ package org.springframework.ai.autoconfigure.stabilityai; -import org.springframework.ai.stabilityai.api.StabilityAiApi; import org.springframework.ai.stabilityai.api.StabilityAiImageOptions; import org.springframework.boot.context.properties.ConfigurationProperties; import org.springframework.boot.context.properties.NestedConfigurationProperty; @@ -30,7 +29,10 @@ public class StabilityAiImageProperties extends StabilityAiParentProperties { public static final String CONFIG_PREFIX = "spring.ai.stabilityai.image"; @NestedConfigurationProperty - private StabilityAiImageOptions options; + private StabilityAiImageOptions options = StabilityAiImageOptions.builder().build(); // stable-diffusion-v1-6 + // is + // default + // model public StabilityAiImageOptions getOptions() { return this.options; diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfigurationIT.java index f4de91be9..936fb748e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfigurationIT.java @@ -16,15 +16,9 @@ package org.springframework.ai.autoconfigure.bedrock.cohere; -import java.util.List; - import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; -import software.amazon.awssdk.regions.Region; - import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.autoconfigure.openai.OpenAiConnectionProperties; -import org.springframework.ai.autoconfigure.openai.OpenAiEmbeddingProperties; import org.springframework.ai.bedrock.cohere.BedrockCohereEmbeddingClient; import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi.CohereEmbeddingModel; import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi.CohereEmbeddingRequest; @@ -32,6 +26,9 @@ import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi.Coher import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import software.amazon.awssdk.regions.Region; + +import java.util.List; import static org.assertj.core.api.Assertions.assertThat; diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiAutoConfigurationIT.java new file mode 100644 index 000000000..7236d343a --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiAutoConfigurationIT.java @@ -0,0 +1,42 @@ +package org.springframework.ai.autoconfigure.stabilityai; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.springframework.ai.image.*; +import org.springframework.ai.stabilityai.StyleEnum; +import org.springframework.ai.stabilityai.api.StabilityAiImageOptions; +import org.springframework.boot.autoconfigure.AutoConfigurations; +import org.springframework.boot.test.context.runner.ApplicationContextRunner; + +import static org.assertj.core.api.Assertions.assertThat; + +@EnabledIfEnvironmentVariable(named = "STABILITYAI_API_KEY", matches = ".*") +public class StabilityAiAutoConfigurationIT { + + private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() + .withPropertyValues("spring.ai.stabilityai.image.api-key=" + System.getenv("STABILITYAI_API_KEY")) + .withConfiguration(AutoConfigurations.of(StabilityAiImageAutoConfiguration.class)); + + @Test + void generate() { + contextRunner.run(context -> { + ImageClient imageClient = context.getBean(ImageClient.class); + StabilityAiImageOptions imageOptions = StabilityAiImageOptions.builder() + .withStylePreset(StyleEnum.PHOTOGRAPHIC) + .build(); + + var instructions = """ + A light cream colored mini golden doodle. + """; + + ImagePrompt imagePrompt = new ImagePrompt(instructions, imageOptions); + ImageResponse imageResponse = imageClient.call(imagePrompt); + + ImageGeneration imageGeneration = imageResponse.getResult(); + Image image = imageGeneration.getOutput(); + + assertThat(image.getB64Json()).isNotEmpty(); + }); + } + +} diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImagePropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImagePropertiesTests.java new file mode 100644 index 000000000..ab704ce83 --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImagePropertiesTests.java @@ -0,0 +1,77 @@ +/* + * 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.autoconfigure.stabilityai; + +import org.junit.jupiter.api.Test; +import org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiAutoConfiguration; +import org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiChatProperties; +import org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiConnectionProperties; +import org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiEmbeddingProperties; +import org.springframework.boot.autoconfigure.AutoConfigurations; +import org.springframework.boot.test.context.runner.ApplicationContextRunner; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * @author Christian Tzolov + * @since 0.8.0 + */ +public class StabilityAiImagePropertiesTests { + + @Test + public void chatPropertiesTest() { + + new ApplicationContextRunner().withPropertyValues( + // @formatter:off + "spring.ai.stabilityai.image.api-key=API_KEY", + "spring.ai.stabilityai.image.base-url=ENDPOINT", + "spring.ai.stabilityai.image.options.n=10", + "spring.ai.stabilityai.image.options.model=MODEL_XYZ", + "spring.ai.stabilityai.image.options.width=512", + "spring.ai.stabilityai.image.options.height=256", + "spring.ai.stabilityai.image.options.response-format=application/json", + "spring.ai.stabilityai.image.options.n=4", + "spring.ai.stabilityai.image.options.cfg-scale=7", + "spring.ai.stabilityai.image.options.clip-guidance-preset=SIMPLE", + "spring.ai.stabilityai.image.options.sampler=K_EULER", + "spring.ai.stabilityai.image.options.seed=0", + "spring.ai.stabilityai.image.options.steps=30", + "spring.ai.stabilityai.image.options.style-preset=neon-punk" + ) + // @formatter:on + .withConfiguration(AutoConfigurations.of(StabilityAiImageAutoConfiguration.class)) + .run(context -> { + var chatProperties = context.getBean(StabilityAiImageProperties.class); + + assertThat(chatProperties.getBaseUrl()).isEqualTo("ENDPOINT"); + assertThat(chatProperties.getApiKey()).isEqualTo("API_KEY"); + assertThat(chatProperties.getOptions().getModel()).isEqualTo("MODEL_XYZ"); + + assertThat(chatProperties.getOptions().getWidth()).isEqualTo(512); + assertThat(chatProperties.getOptions().getHeight()).isEqualTo(256); + assertThat(chatProperties.getOptions().getResponseFormat()).isEqualTo("application/json"); + assertThat(chatProperties.getOptions().getN()).isEqualTo(4); + assertThat(chatProperties.getOptions().getCfgScale()).isEqualTo(7); + assertThat(chatProperties.getOptions().getClipGuidancePreset()).isEqualTo("SIMPLE"); + assertThat(chatProperties.getOptions().getSampler()).isEqualTo("K_EULER"); + assertThat(chatProperties.getOptions().getSeed()).isEqualTo(0); + assertThat(chatProperties.getOptions().getSteps()).isEqualTo(30); + assertThat(chatProperties.getOptions().getStylePreset()).isEqualTo("neon-punk"); + }); + } + +}