Streamline ImageOptions

* Add style to option abstraction
* Use abstraction in Observations directly instead of dedicated implementation
* Clean-up the merge of runtime and default image options in OpenAI and Stability AI

Related to #1148

Signed-off-by: Thomas Vitale <ThomasVitale@users.noreply.github.com>
This commit is contained in:
Thomas Vitale
2024-08-03 00:26:16 +02:00
committed by Christian Tzolov
parent 80007d4d6c
commit 17ba1fc3ba
17 changed files with 194 additions and 382 deletions

View File

@@ -11,6 +11,7 @@ import com.fasterxml.jackson.annotation.JsonProperty;
* The configuration information for a image generation request.
*
* @author Benoit Moussaud
* @author Thomas Vitale
* @since 1.0.0 M1
*/
@JsonInclude(JsonInclude.Include.NON_NULL)
@@ -88,10 +89,15 @@ public class AzureOpenAiImageOptions implements ImageOptions {
@JsonProperty("user")
private String user;
@Override
public Integer getN() {
return n;
}
public void setN(Integer n) {
this.n = n;
}
@Override
public String getModel() {
return model;
@@ -101,10 +107,7 @@ public class AzureOpenAiImageOptions implements ImageOptions {
this.model = model;
}
public void setN(Integer n) {
this.n = n;
}
@Override
public Integer getWidth() {
return width;
}
@@ -114,6 +117,7 @@ public class AzureOpenAiImageOptions implements ImageOptions {
this.size = this.width + "x" + this.height;
}
@Override
public Integer getHeight() {
return height;
}
@@ -123,6 +127,7 @@ public class AzureOpenAiImageOptions implements ImageOptions {
this.size = this.width + "x" + this.height;
}
@Override
public String getResponseFormat() {
return responseFormat;
}
@@ -158,6 +163,7 @@ public class AzureOpenAiImageOptions implements ImageOptions {
this.quality = quality;
}
@Override
public String getStyle() {
return style;
}

View File

@@ -29,7 +29,6 @@ import org.springframework.ai.image.observation.DefaultImageModelObservationConv
import org.springframework.ai.image.observation.ImageModelObservationConvention;
import org.springframework.ai.image.observation.ImageModelObservationContext;
import org.springframework.ai.image.observation.ImageModelObservationDocumentation;
import org.springframework.ai.image.observation.ImageModelRequestOptions;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.observation.AiOperationMetadata;
import org.springframework.ai.observation.conventions.AiOperationType;
@@ -40,7 +39,6 @@ import org.springframework.ai.retry.RetryUtils;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
import java.util.List;
@@ -128,12 +126,13 @@ public class OpenAiImageModel implements ImageModel {
@Override
public ImageResponse call(ImagePrompt imagePrompt) {
OpenAiImageApi.OpenAiImageRequest imageRequest = createRequest(imagePrompt);
OpenAiImageOptions requestImageOptions = mergeOptions(imagePrompt.getOptions(), this.defaultOptions);
OpenAiImageApi.OpenAiImageRequest imageRequest = createRequest(imagePrompt, requestImageOptions);
var observationContext = ImageModelObservationContext.builder()
.imagePrompt(imagePrompt)
.operationMetadata(buildOperationMetadata())
.requestOptions(buildRequestOptions(imageRequest))
.requestOptions(requestImageOptions)
.build();
return ImageModelObservationDocumentation.IMAGE_MODEL_OPERATION
@@ -151,23 +150,14 @@ public class OpenAiImageModel implements ImageModel {
});
}
private OpenAiImageApi.OpenAiImageRequest createRequest(ImagePrompt imagePrompt) {
private OpenAiImageApi.OpenAiImageRequest createRequest(ImagePrompt imagePrompt,
OpenAiImageOptions requestImageOptions) {
String instructions = imagePrompt.getInstructions().get(0).getText();
OpenAiImageApi.OpenAiImageRequest imageRequest = new OpenAiImageApi.OpenAiImageRequest(instructions,
OpenAiImageApi.DEFAULT_IMAGE_MODEL);
if (this.defaultOptions != null) {
imageRequest = ModelOptionsUtils.merge(this.defaultOptions, imageRequest,
OpenAiImageApi.OpenAiImageRequest.class);
}
if (imagePrompt.getOptions() != null) {
imageRequest = ModelOptionsUtils.merge(toOpenAiImageOptions(imagePrompt.getOptions()), imageRequest,
OpenAiImageApi.OpenAiImageRequest.class);
}
return imageRequest;
return ModelOptionsUtils.merge(requestImageOptions, imageRequest, OpenAiImageApi.OpenAiImageRequest.class);
}
private ImageResponse convertResponse(ResponseEntity<OpenAiImageApi.OpenAiImageResponse> imageResponseEntity,
@@ -188,44 +178,27 @@ public class OpenAiImageModel implements ImageModel {
}
/**
* Convert the {@link ImageOptions} into {@link OpenAiImageOptions}.
* @param runtimeImageOptions the image options to use.
* @return the converted {@link OpenAiImageOptions}.
* Merge runtime and default {@link ImageOptions} to compute the final options to use
* in the request.
*/
private OpenAiImageOptions toOpenAiImageOptions(ImageOptions runtimeImageOptions) {
OpenAiImageOptions.Builder openAiImageOptionsBuilder = OpenAiImageOptions.builder();
if (runtimeImageOptions != null) {
// Handle portable image options
if (runtimeImageOptions.getN() != null) {
openAiImageOptionsBuilder.withN(runtimeImageOptions.getN());
}
if (runtimeImageOptions.getModel() != null) {
openAiImageOptionsBuilder.withModel(runtimeImageOptions.getModel());
}
if (runtimeImageOptions.getResponseFormat() != null) {
openAiImageOptionsBuilder.withResponseFormat(runtimeImageOptions.getResponseFormat());
}
if (runtimeImageOptions.getWidth() != null) {
openAiImageOptionsBuilder.withWidth(runtimeImageOptions.getWidth());
}
if (runtimeImageOptions.getHeight() != null) {
openAiImageOptionsBuilder.withHeight(runtimeImageOptions.getHeight());
}
// Handle OpenAI specific image options
if (runtimeImageOptions instanceof OpenAiImageOptions) {
OpenAiImageOptions runtimeOpenAiImageOptions = (OpenAiImageOptions) runtimeImageOptions;
if (runtimeOpenAiImageOptions.getQuality() != null) {
openAiImageOptionsBuilder.withQuality(runtimeOpenAiImageOptions.getQuality());
}
if (runtimeOpenAiImageOptions.getStyle() != null) {
openAiImageOptionsBuilder.withStyle(runtimeOpenAiImageOptions.getStyle());
}
if (runtimeOpenAiImageOptions.getUser() != null) {
openAiImageOptionsBuilder.withUser(runtimeOpenAiImageOptions.getUser());
}
}
private OpenAiImageOptions mergeOptions(ImageOptions runtimeOptions, OpenAiImageOptions defaultOptions) {
if (runtimeOptions == null) {
return defaultOptions;
}
return openAiImageOptionsBuilder.build();
return OpenAiImageOptions.builder()
// Handle portable image options
.withModel(ModelOptionsUtils.mergeOption(runtimeOptions.getModel(), defaultOptions.getModel()))
.withN(ModelOptionsUtils.mergeOption(runtimeOptions.getN(), defaultOptions.getN()))
.withResponseFormat(ModelOptionsUtils.mergeOption(runtimeOptions.getResponseFormat(),
defaultOptions.getResponseFormat()))
.withWidth(ModelOptionsUtils.mergeOption(runtimeOptions.getWidth(), defaultOptions.getWidth()))
.withHeight(ModelOptionsUtils.mergeOption(runtimeOptions.getHeight(), defaultOptions.getHeight()))
.withStyle(ModelOptionsUtils.mergeOption(runtimeOptions.getStyle(), defaultOptions.getStyle()))
// Handle OpenAI specific image options
.withQuality(defaultOptions.getQuality())
.withUser(defaultOptions.getUser())
.build();
}
private AiOperationMetadata buildOperationMetadata() {
@@ -235,17 +208,6 @@ public class OpenAiImageModel implements ImageModel {
.build();
}
private ImageModelRequestOptions buildRequestOptions(OpenAiImageApi.OpenAiImageRequest request) {
return ImageModelRequestOptions.builder()
.model(StringUtils.hasText(request.model()) ? request.model() : "unknown")
.n(request.n())
.width(request.size() != null ? Integer.parseInt(request.size().split("x")[0]) : null)
.height(request.size() != null ? Integer.parseInt(request.size().split("x")[1]) : null)
.responseFormat(request.responseFormat())
.style(request.style())
.build();
}
/**
* Use the provided convention for reporting observation data
* @param observationConvention The provided convention

View File

@@ -208,6 +208,7 @@ public class OpenAiImageOptions implements ImageOptions {
this.size = this.width + "x" + this.height;
}
@Override
public String getStyle() {
return this.style;
}

View File

@@ -164,16 +164,17 @@ public class QianFanImageOptions implements ImageOptions {
return this.height;
}
@Override
public String getResponseFormat() {
return null;
}
public void setHeight(Integer height) {
this.height = height;
this.size = this.width + "x" + this.height;
}
@Override
public String getResponseFormat() {
return null;
}
@Override
public String getStyle() {
return this.style;
}
@@ -190,18 +191,17 @@ public class QianFanImageOptions implements ImageOptions {
this.user = user;
}
public void setSize(String size) {
this.size = size;
}
public String getSize() {
if (this.size != null) {
return this.size;
}
return (this.width != null && this.height != null) ? this.width + "x" + this.height : null;
}
public void setSize(String size) {
this.size = size;
}
@Override
public boolean equals(Object o) {
if (this == o)

View File

@@ -18,9 +18,6 @@ package org.springframework.ai.stabilityai;
import java.util.List;
import java.util.stream.Collectors;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.image.Image;
import org.springframework.ai.image.ImageModel;
import org.springframework.ai.image.ImageGeneration;
@@ -39,9 +36,7 @@ import org.springframework.util.Assert;
*/
public class StabilityAiImageModel implements ImageModel {
private final Logger logger = LoggerFactory.getLogger(getClass());
private StabilityAiImageOptions options;
private final StabilityAiImageOptions defaultOptions;
private final StabilityAiApi stabilityAiApi;
@@ -49,15 +44,15 @@ public class StabilityAiImageModel implements ImageModel {
this(stabilityAiApi, StabilityAiImageOptions.builder().build());
}
public StabilityAiImageModel(StabilityAiApi stabilityAiApi, StabilityAiImageOptions options) {
public StabilityAiImageModel(StabilityAiApi stabilityAiApi, StabilityAiImageOptions defaultOptions) {
Assert.notNull(stabilityAiApi, "StabilityAiApi must not be null");
Assert.notNull(options, "StabilityAiImageOptions must not be null");
Assert.notNull(defaultOptions, "StabilityAiImageOptions must not be null");
this.stabilityAiApi = stabilityAiApi;
this.options = options;
this.defaultOptions = defaultOptions;
}
public StabilityAiImageOptions getOptions() {
return this.options;
return this.defaultOptions;
}
/**
@@ -69,18 +64,15 @@ public class StabilityAiImageModel implements ImageModel {
* @return the ImageResponse generated by the StabilityAiImageModel
*/
public ImageResponse call(ImagePrompt imagePrompt) {
ImageOptions runtimeOptions = imagePrompt.getOptions();
// Merge the runtime options passed via the prompt with the StabilityAiImageModel
// options configured via Autoconfiguration.
// Merge the runtime options passed via the prompt with the default options
// configured via the constructor.
// Runtime options overwrite StabilityAiImageModel options
StabilityAiImageOptions optionsToUse = ModelOptionsUtils.merge(runtimeOptions, this.options,
StabilityAiImageOptions.class);
StabilityAiImageOptions requestImageOptions = mergeOptions(imagePrompt.getOptions(), this.defaultOptions);
// Copy the org.springframework.ai.model derived ImagePrompt and ImageOptions data
// types to the data types used in StabilityAiApi
StabilityAiApi.GenerateImageRequest generateImageRequest = getGenerateImageRequest(imagePrompt, optionsToUse);
StabilityAiApi.GenerateImageRequest generateImageRequest = getGenerateImageRequest(imagePrompt,
requestImageOptions);
// Make the request
StabilityAiApi.GenerateImageResponse generateImageResponse = this.stabilityAiApi
@@ -88,13 +80,11 @@ public class StabilityAiImageModel implements ImageModel {
// Convert to org.springframework.ai.model derived ImageResponse data type
return convertResponse(generateImageResponse);
}
private static StabilityAiApi.GenerateImageRequest getGenerateImageRequest(ImagePrompt stabilityAiImagePrompt,
StabilityAiImageOptions optionsToUse) {
StabilityAiApi.GenerateImageRequest.Builder builder = new StabilityAiApi.GenerateImageRequest.Builder();
StabilityAiApi.GenerateImageRequest generateImageRequest = builder
return new StabilityAiApi.GenerateImageRequest.Builder()
.withTextPrompts(stabilityAiImagePrompt.getInstructions()
.stream()
.map(message -> new StabilityAiApi.GenerateImageRequest.TextPrompts(message.getText(),
@@ -110,60 +100,44 @@ public class StabilityAiImageModel implements ImageModel {
.withSteps(optionsToUse.getSteps())
.withStylePreset(optionsToUse.getStylePreset())
.build();
return generateImageRequest;
}
private ImageResponse convertResponse(StabilityAiApi.GenerateImageResponse generateImageResponse) {
List<ImageGeneration> imageGenerationList = generateImageResponse.artifacts().stream().map(entry -> {
return new ImageGeneration(new Image(null, entry.base64()),
new StabilityAiImageGenerationMetadata(entry.finishReason(), entry.seed()));
}).toList();
List<ImageGeneration> imageGenerationList = generateImageResponse.artifacts()
.stream()
.map(entry -> new ImageGeneration(new Image(null, entry.base64()),
new StabilityAiImageGenerationMetadata(entry.finishReason(), entry.seed())))
.toList();
return new ImageResponse(imageGenerationList, new ImageResponseMetadata());
}
private StabilityAiImageOptions convertOptions(ImageOptions runtimeOptions) {
StabilityAiImageOptions.Builder builder = StabilityAiImageOptions.builder();
/**
* Merge runtime and default {@link ImageOptions} to compute the final options to use
* in the request.
*/
private StabilityAiImageOptions mergeOptions(ImageOptions runtimeOptions, StabilityAiImageOptions defaultOptions) {
if (runtimeOptions == null) {
return builder.build();
return defaultOptions;
}
if (runtimeOptions.getN() != null) {
builder.withN(runtimeOptions.getN());
}
if (runtimeOptions.getModel() != null) {
builder.withModel(runtimeOptions.getModel());
}
if (runtimeOptions.getResponseFormat() != null) {
builder.withResponseFormat(runtimeOptions.getResponseFormat());
}
if (runtimeOptions.getWidth() != null) {
builder.withWidth(runtimeOptions.getWidth());
}
if (runtimeOptions.getHeight() != null) {
builder.withHeight(runtimeOptions.getHeight());
}
if (runtimeOptions instanceof StabilityAiImageOptions) {
StabilityAiImageOptions stabilityAiImageOptions = (StabilityAiImageOptions) runtimeOptions;
if (stabilityAiImageOptions.getCfgScale() != null) {
builder.withCfgScale(stabilityAiImageOptions.getCfgScale());
}
if (stabilityAiImageOptions.getClipGuidancePreset() != null) {
builder.withClipGuidancePreset(stabilityAiImageOptions.getClipGuidancePreset());
}
if (stabilityAiImageOptions.getSampler() != null) {
builder.withSampler(stabilityAiImageOptions.getSampler());
}
if (stabilityAiImageOptions.getSeed() != null) {
builder.withSeed(stabilityAiImageOptions.getSeed());
}
if (stabilityAiImageOptions.getSteps() != null) {
builder.withSteps(stabilityAiImageOptions.getSteps());
}
if (stabilityAiImageOptions.getStylePreset() != null) {
builder.withStylePreset(stabilityAiImageOptions.getStylePreset());
}
}
return builder.build();
return StabilityAiImageOptions.builder()
// Handle portable image options
.withModel(ModelOptionsUtils.mergeOption(runtimeOptions.getModel(), defaultOptions.getModel()))
.withN(ModelOptionsUtils.mergeOption(runtimeOptions.getN(), defaultOptions.getN()))
.withResponseFormat(ModelOptionsUtils.mergeOption(runtimeOptions.getResponseFormat(),
defaultOptions.getResponseFormat()))
.withWidth(ModelOptionsUtils.mergeOption(runtimeOptions.getWidth(), defaultOptions.getWidth()))
.withHeight(ModelOptionsUtils.mergeOption(runtimeOptions.getHeight(), defaultOptions.getHeight()))
.withStylePreset(ModelOptionsUtils.mergeOption(runtimeOptions.getStyle(), defaultOptions.getStyle()))
// Handle Stability AI specific image options
.withCfgScale(defaultOptions.getCfgScale())
.withClipGuidancePreset(defaultOptions.getClipGuidancePreset())
.withSampler(defaultOptions.getSampler())
.withSeed(defaultOptions.getSeed())
.withSteps(defaultOptions.getSteps())
.withStylePreset(defaultOptions.getStylePreset())
.build();
}
}

View File

@@ -451,6 +451,11 @@ public class StabilityAiImageOptions implements ImageOptions {
this.steps = steps;
}
@Override
public String getStyle() {
return getStylePreset();
}
public String getStylePreset() {
return stylePreset;
}

View File

@@ -97,6 +97,10 @@ public class ZhiPuAiImageOptions implements ImageOptions {
return this.model;
}
public void setModel(String model) {
this.model = model;
}
@Override
public Integer getWidth() {
return null;
@@ -112,8 +116,9 @@ public class ZhiPuAiImageOptions implements ImageOptions {
return null;
}
public void setModel(String model) {
this.model = model;
@Override
public String getStyle() {
return null;
}
public String getUser() {

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.image;
import org.springframework.ai.model.ModelOptions;
import org.springframework.lang.Nullable;
/**
* ImageOptions represent the common options, portable across different image generation
@@ -23,14 +24,22 @@ import org.springframework.ai.model.ModelOptions;
*/
public interface ImageOptions extends ModelOptions {
@Nullable
Integer getN();
@Nullable
String getModel();
@Nullable
Integer getWidth();
@Nullable
Integer getHeight();
String getResponseFormat(); // openai - url or base64 : stability ai byte[] or base64
@Nullable
String getResponseFormat();
@Nullable
String getStyle();
}

View File

@@ -17,7 +17,7 @@ package org.springframework.ai.image;
public class ImageOptionsBuilder {
private class ImageModelOptionsImpl implements ImageOptions {
private static class DefaultImageModelOptions implements ImageOptions {
private Integer n;
@@ -29,6 +29,8 @@ public class ImageOptionsBuilder {
private String responseFormat;
private String style;
@Override
public Integer getN() {
return n;
@@ -74,9 +76,18 @@ public class ImageOptionsBuilder {
this.height = height;
}
@Override
public String getStyle() {
return style;
}
public void setStyle(String style) {
this.style = style;
}
}
private final ImageModelOptionsImpl options = new ImageModelOptionsImpl();
private final DefaultImageModelOptions options = new DefaultImageModelOptions();
private ImageOptionsBuilder() {
@@ -111,6 +122,11 @@ public class ImageOptionsBuilder {
return this;
}
public ImageOptionsBuilder withStyle(String style) {
options.setStyle(style);
return this;
}
public ImageOptions build() {
return options;
}

View File

@@ -27,6 +27,9 @@ import org.springframework.util.StringUtils;
*/
public class DefaultImageModelObservationConvention implements ImageModelObservationConvention {
private static final KeyValue REQUEST_MODEL_NONE = KeyValue
.of(ImageModelObservationDocumentation.LowCardinalityKeyNames.REQUEST_MODEL, KeyValue.NONE_VALUE);
private static final KeyValue REQUEST_IMAGE_RESPONSE_FORMAT_NONE = KeyValue.of(
ImageModelObservationDocumentation.HighCardinalityKeyNames.REQUEST_IMAGE_RESPONSE_FORMAT,
KeyValue.NONE_VALUE);
@@ -46,8 +49,11 @@ public class DefaultImageModelObservationConvention implements ImageModelObserva
@Override
public String getContextualName(ImageModelObservationContext context) {
return "%s %s".formatted(context.getOperationMetadata().operationType(),
context.getRequestOptions().getModel());
if (StringUtils.hasText(context.getRequestOptions().getModel())) {
return "%s %s".formatted(context.getOperationMetadata().operationType(),
context.getRequestOptions().getModel());
}
return context.getOperationMetadata().operationType();
}
@Override
@@ -57,7 +63,7 @@ public class DefaultImageModelObservationConvention implements ImageModelObserva
protected KeyValue aiOperationType(ImageModelObservationContext context) {
return KeyValue.of(ImageModelObservationDocumentation.LowCardinalityKeyNames.AI_OPERATION_TYPE,
context.getOperationMetadata().operationType());
context.getOperationType());
}
protected KeyValue aiProvider(ImageModelObservationContext context) {
@@ -66,8 +72,11 @@ public class DefaultImageModelObservationConvention implements ImageModelObserva
}
protected KeyValue requestModel(ImageModelObservationContext context) {
return KeyValue.of(ImageModelObservationDocumentation.LowCardinalityKeyNames.REQUEST_MODEL,
context.getRequestOptions().getModel());
if (StringUtils.hasText(context.getRequestOptions().getModel())) {
return KeyValue.of(ImageModelObservationDocumentation.LowCardinalityKeyNames.REQUEST_MODEL,
context.getRequestOptions().getModel());
}
return REQUEST_MODEL_NONE;
}
@Override

View File

@@ -15,10 +15,12 @@
*/
package org.springframework.ai.image.observation;
import org.springframework.ai.image.ImageOptions;
import org.springframework.ai.image.ImagePrompt;
import org.springframework.ai.image.ImageResponse;
import org.springframework.ai.model.observation.ModelObservationContext;
import org.springframework.ai.observation.AiOperationMetadata;
import org.springframework.ai.observation.conventions.AiOperationType;
import org.springframework.util.Assert;
/**
@@ -29,19 +31,23 @@ import org.springframework.util.Assert;
*/
public class ImageModelObservationContext extends ModelObservationContext<ImagePrompt, ImageResponse> {
private final ImageModelRequestOptions requestOptions;
private final ImageOptions requestOptions;
ImageModelObservationContext(ImagePrompt imagePrompt, AiOperationMetadata operationMetadata,
ImageModelRequestOptions requestOptions) {
ImageOptions requestOptions) {
super(imagePrompt, operationMetadata);
Assert.notNull(requestOptions, "requestOptions cannot be null");
this.requestOptions = requestOptions;
}
public ImageModelRequestOptions getRequestOptions() {
public ImageOptions getRequestOptions() {
return requestOptions;
}
public String getOperationType() {
return AiOperationType.IMAGE.value();
}
public static Builder builder() {
return new Builder();
}
@@ -52,7 +58,7 @@ public class ImageModelObservationContext extends ModelObservationContext<ImageP
private AiOperationMetadata operationMetadata;
private ImageModelRequestOptions requestOptions;
private ImageOptions requestOptions;
private Builder() {
}
@@ -67,7 +73,7 @@ public class ImageModelObservationContext extends ModelObservationContext<ImageP
return this;
}
public Builder requestOptions(ImageModelRequestOptions requestOptions) {
public Builder requestOptions(ImageOptions requestOptions) {
this.requestOptions = requestOptions;
return this;
}

View File

@@ -1,154 +0,0 @@
/*
* Copyright 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.image.observation;
import org.springframework.ai.image.ImageOptions;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
/**
* Represents client-side options for image model requests.
*
* @author Thomas Vitale
* @since 1.0.0
*/
public class ImageModelRequestOptions implements ImageOptions {
private final String model;
@Nullable
private final Integer n;
@Nullable
private final Integer width;
@Nullable
private final Integer height;
@Nullable
private final String responseFormat;
@Nullable
private final String style;
ImageModelRequestOptions(Builder builder) {
Assert.hasText(builder.model, "model cannot be null or empty");
this.model = builder.model;
this.n = builder.n;
this.width = builder.width;
this.height = builder.height;
this.responseFormat = builder.responseFormat;
this.style = builder.style;
}
public static Builder builder() {
return new Builder();
}
public static class Builder {
private String model;
@Nullable
private Integer n;
@Nullable
private Integer width;
@Nullable
private Integer height;
@Nullable
private String responseFormat;
@Nullable
private String style;
private Builder() {
}
public Builder model(String model) {
this.model = model;
return this;
}
public Builder n(@Nullable Integer n) {
this.n = n;
return this;
}
public Builder width(@Nullable Integer width) {
this.width = width;
return this;
}
public Builder height(@Nullable Integer height) {
this.height = height;
return this;
}
public Builder responseFormat(@Nullable String responseFormat) {
this.responseFormat = responseFormat;
return this;
}
public Builder style(@Nullable String style) {
this.style = style;
return this;
}
public ImageModelRequestOptions build() {
return new ImageModelRequestOptions(this);
}
}
@Override
public String getModel() {
return model;
}
@Override
@Nullable
public Integer getN() {
return n;
}
@Override
@Nullable
public Integer getWidth() {
return width;
}
@Override
@Nullable
public Integer getHeight() {
return height;
}
@Override
@Nullable
public String getResponseFormat() {
return responseFormat;
}
@Nullable
public String getStyle() {
return style;
}
}

View File

@@ -50,11 +50,13 @@ import org.springframework.beans.BeanWrapper;
import org.springframework.beans.BeanWrapperImpl;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.ObjectUtils;
/**
* Utility class for manipulating {@link ModelOptions} objects.
*
* @author Christian Tzolov
* @author Thomas Vitale
* @since 0.8.0
*/
public abstract class ModelOptionsUtils {
@@ -391,4 +393,11 @@ public abstract class ModelOptionsUtils {
}
}
/**
* Return the runtime value if not empty, or else the default value.
*/
public static <T> T mergeOption(T runtimeValue, T defaultValue) {
return ObjectUtils.isEmpty(runtimeValue) ? defaultValue : runtimeValue;
}
}

View File

@@ -18,6 +18,7 @@ package org.springframework.ai.image.observation;
import io.micrometer.common.KeyValue;
import io.micrometer.observation.Observation;
import org.junit.jupiter.api.Test;
import org.springframework.ai.image.ImageOptionsBuilder;
import org.springframework.ai.image.ImagePrompt;
import org.springframework.ai.observation.AiOperationMetadata;
import org.springframework.ai.observation.conventions.AiObservationAttributes;
@@ -41,32 +42,42 @@ class DefaultImageModelObservationConventionTests {
}
@Test
void shouldHaveContextualName() {
void contextualNameWhenModelIsDefined() {
ImageModelObservationContext observationContext = ImageModelObservationContext.builder()
.imagePrompt(generateImagePrompt())
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder().model("mistral").build())
.requestOptions(ImageOptionsBuilder.builder().withModel("mistral").build())
.build();
assertThat(this.observationConvention.getContextualName(observationContext)).isEqualTo("image mistral");
}
@Test
void contextualNameWhenModelIsNotDefined() {
ImageModelObservationContext observationContext = ImageModelObservationContext.builder()
.imagePrompt(generateImagePrompt())
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageOptionsBuilder.builder().build())
.build();
assertThat(this.observationConvention.getContextualName(observationContext)).isEqualTo("image");
}
@Test
void supportsOnlyImageModelObservationContext() {
ImageModelObservationContext observationContext = ImageModelObservationContext.builder()
.imagePrompt(generateImagePrompt())
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder().model("mistral").build())
.requestOptions(ImageOptionsBuilder.builder().withModel("mistral").build())
.build();
assertThat(this.observationConvention.supportsContext(observationContext)).isTrue();
assertThat(this.observationConvention.supportsContext(new Observation.Context())).isFalse();
}
@Test
void shouldHaveRequiredLowCardinalityKeyValues() {
void shouldHaveLowCardinalityKeyValuesWhenDefined() {
ImageModelObservationContext observationContext = ImageModelObservationContext.builder()
.imagePrompt(generateImagePrompt())
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder().model("mistral").build())
.requestOptions(ImageOptionsBuilder.builder().withModel("mistral").build())
.build();
assertThat(this.observationConvention.getLowCardinalityKeyValues(observationContext)).contains(
KeyValue.of(AiObservationAttributes.AI_OPERATION_TYPE.value(), "image"),
@@ -75,17 +86,17 @@ class DefaultImageModelObservationConventionTests {
}
@Test
void shouldHaveOptionalHighCardinalityKeyValues() {
void shouldHaveHighCardinalityKeyValuesWhenDefined() {
ImageModelObservationContext observationContext = ImageModelObservationContext.builder()
.imagePrompt(generateImagePrompt())
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder()
.model("mistral")
.n(1)
.height(1080)
.width(1920)
.style("sketch")
.responseFormat("base64")
.requestOptions(ImageOptionsBuilder.builder()
.withModel("mistral")
.withN(1)
.withHeight(1080)
.withWidth(1920)
.withStyle("sketch")
.withResponseFormat("base64")
.build())
.build();
@@ -96,13 +107,15 @@ class DefaultImageModelObservationConventionTests {
}
@Test
void shouldHaveMissingHighCardinalityKeyValues() {
void shouldHaveNoneKeyValuesWhenMissing() {
ImageModelObservationContext observationContext = ImageModelObservationContext.builder()
.imagePrompt(generateImagePrompt())
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder().model("mistral").build())
.requestOptions(ImageOptionsBuilder.builder().build())
.build();
assertThat(this.observationConvention.getLowCardinalityKeyValues(observationContext))
.contains(KeyValue.of(AiObservationAttributes.REQUEST_MODEL.value(), KeyValue.NONE_VALUE));
assertThat(this.observationConvention.getHighCardinalityKeyValues(observationContext)).contains(
KeyValue.of(AiObservationAttributes.REQUEST_IMAGE_RESPONSE_FORMAT.value(), KeyValue.NONE_VALUE),
KeyValue.of(AiObservationAttributes.REQUEST_IMAGE_SIZE.value(), KeyValue.NONE_VALUE),

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.image.observation;
import org.junit.jupiter.api.Test;
import org.springframework.ai.image.ImageOptionsBuilder;
import org.springframework.ai.image.ImagePrompt;
import org.springframework.ai.observation.AiOperationMetadata;
import org.springframework.ai.observation.conventions.AiOperationType;
@@ -36,7 +37,7 @@ class ImageModelObservationContextTests {
var observationContext = ImageModelObservationContext.builder()
.imagePrompt(generateImagePrompt())
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder().model("supersun").build())
.requestOptions(ImageOptionsBuilder.builder().withModel("supersun").build())
.build();
assertThat(observationContext).isNotNull();

View File

@@ -19,6 +19,7 @@ import io.micrometer.common.KeyValue;
import io.micrometer.observation.Observation;
import org.junit.jupiter.api.Test;
import org.springframework.ai.image.ImageMessage;
import org.springframework.ai.image.ImageOptionsBuilder;
import org.springframework.ai.image.ImagePrompt;
import org.springframework.ai.observation.AiOperationMetadata;
import org.springframework.ai.observation.conventions.AiObservationAttributes;
@@ -51,7 +52,7 @@ class ImageModelPromptContentObservationFilterTests {
var expectedContext = ImageModelObservationContext.builder()
.imagePrompt(new ImagePrompt(""))
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder().model("mistral").build())
.requestOptions(ImageOptionsBuilder.builder().withModel("mistral").build())
.build();
var actualContext = observationFilter.map(expectedContext);
@@ -63,7 +64,7 @@ class ImageModelPromptContentObservationFilterTests {
var originalContext = ImageModelObservationContext.builder()
.imagePrompt(new ImagePrompt("supercalifragilisticexpialidocious"))
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder().model("mistral").build())
.requestOptions(ImageOptionsBuilder.builder().withModel("mistral").build())
.build();
var augmentedContext = observationFilter.map(originalContext);
@@ -77,7 +78,7 @@ class ImageModelPromptContentObservationFilterTests {
.imagePrompt(new ImagePrompt(List.of(new ImageMessage("you're a chimney sweep"),
new ImageMessage("supercalifragilisticexpialidocious"))))
.operationMetadata(generateOperationMetadata())
.requestOptions(ImageModelRequestOptions.builder().model("mistral").build())
.requestOptions(ImageOptionsBuilder.builder().withModel("mistral").build())
.build();
var augmentedContext = observationFilter.map(originalContext);

View File

@@ -1,51 +0,0 @@
/*
* Copyright 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.image.observation;
import org.junit.jupiter.api.Test;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
/**
* Unit tests for {@link ImageModelRequestOptions}.
*
* @author Thomas Vitale
*/
class ImageModelRequestOptionsTests {
@Test
void whenMandatoryRequestOptionsThenReturn() {
var requestOptions = ImageModelRequestOptions.builder().model("rowena").build();
assertThat(requestOptions).isNotNull();
}
@Test
void whenModelIsNullThenThrow() {
assertThatThrownBy(() -> ImageModelRequestOptions.builder().build())
.isInstanceOf(IllegalArgumentException.class)
.hasMessageContaining("model cannot be null or empty");
}
@Test
void whenModelIsEmptyThenThrow() {
assertThatThrownBy(() -> ImageModelRequestOptions.builder().model("").build())
.isInstanceOf(IllegalArgumentException.class)
.hasMessageContaining("model cannot be null or empty");
}
}