From 17ba1fc3baea902e983cc5ce9b5e2b93a18badcc Mon Sep 17 00:00:00 2001 From: Thomas Vitale Date: Sat, 3 Aug 2024 00:26:16 +0200 Subject: [PATCH] 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 --- .../azure/openai/AzureOpenAiImageOptions.java | 14 +- .../ai/openai/OpenAiImageModel.java | 88 +++------- .../ai/openai/OpenAiImageOptions.java | 1 + .../ai/qianfan/QianFanImageOptions.java | 20 +-- .../ai/stabilityai/StabilityAiImageModel.java | 106 +++++------- .../api/StabilityAiImageOptions.java | 5 + .../ai/zhipuai/ZhiPuAiImageOptions.java | 9 +- .../ai/image/ImageOptions.java | 11 +- .../ai/image/ImageOptionsBuilder.java | 20 ++- ...efaultImageModelObservationConvention.java | 19 ++- .../ImageModelObservationContext.java | 16 +- .../observation/ImageModelRequestOptions.java | 154 ------------------ .../ai/model/ModelOptionsUtils.java | 9 + ...tImageModelObservationConventionTests.java | 43 +++-- .../ImageModelObservationContextTests.java | 3 +- ...elPromptContentObservationFilterTests.java | 7 +- .../ImageModelRequestOptionsTests.java | 51 ------ 17 files changed, 194 insertions(+), 382 deletions(-) delete mode 100644 spring-ai-core/src/main/java/org/springframework/ai/image/observation/ImageModelRequestOptions.java delete mode 100644 spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelRequestOptionsTests.java diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiImageOptions.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiImageOptions.java index 064551556..be15fbfd1 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiImageOptions.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiImageOptions.java @@ -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; } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java index 43b88cf07..88e6656d6 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java @@ -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 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 diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageOptions.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageOptions.java index cd2f74e2c..9b7e30e82 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageOptions.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageOptions.java @@ -208,6 +208,7 @@ public class OpenAiImageOptions implements ImageOptions { this.size = this.width + "x" + this.height; } + @Override public String getStyle() { return this.style; } diff --git a/models/spring-ai-qianfan/src/main/java/org/springframework/ai/qianfan/QianFanImageOptions.java b/models/spring-ai-qianfan/src/main/java/org/springframework/ai/qianfan/QianFanImageOptions.java index cf01354f6..bd4d3bee9 100644 --- a/models/spring-ai-qianfan/src/main/java/org/springframework/ai/qianfan/QianFanImageOptions.java +++ b/models/spring-ai-qianfan/src/main/java/org/springframework/ai/qianfan/QianFanImageOptions.java @@ -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) diff --git a/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageModel.java b/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageModel.java index bf6f41f1a..e1db5e2ac 100644 --- a/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageModel.java +++ b/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageModel.java @@ -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 imageGenerationList = generateImageResponse.artifacts().stream().map(entry -> { - return new ImageGeneration(new Image(null, entry.base64()), - new StabilityAiImageGenerationMetadata(entry.finishReason(), entry.seed())); - }).toList(); + List 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(); } } diff --git a/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/api/StabilityAiImageOptions.java b/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/api/StabilityAiImageOptions.java index 3e32017c5..a6cccaf8e 100644 --- a/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/api/StabilityAiImageOptions.java +++ b/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/api/StabilityAiImageOptions.java @@ -451,6 +451,11 @@ public class StabilityAiImageOptions implements ImageOptions { this.steps = steps; } + @Override + public String getStyle() { + return getStylePreset(); + } + public String getStylePreset() { return stylePreset; } diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageOptions.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageOptions.java index 357eb536a..1f80d1a74 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageOptions.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageOptions.java @@ -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() { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageOptions.java b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageOptions.java index 8914862cb..af8962029 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageOptions.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageOptions.java @@ -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(); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageOptionsBuilder.java b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageOptionsBuilder.java index c7b7f3eb1..7917f3df5 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageOptionsBuilder.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageOptionsBuilder.java @@ -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; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/image/observation/DefaultImageModelObservationConvention.java b/spring-ai-core/src/main/java/org/springframework/ai/image/observation/DefaultImageModelObservationConvention.java index 06979d566..e7adc23da 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/image/observation/DefaultImageModelObservationConvention.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/image/observation/DefaultImageModelObservationConvention.java @@ -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 diff --git a/spring-ai-core/src/main/java/org/springframework/ai/image/observation/ImageModelObservationContext.java b/spring-ai-core/src/main/java/org/springframework/ai/image/observation/ImageModelObservationContext.java index 3a9c18e92..599804506 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/image/observation/ImageModelObservationContext.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/image/observation/ImageModelObservationContext.java @@ -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 { - 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 T mergeOption(T runtimeValue, T defaultValue) { + return ObjectUtils.isEmpty(runtimeValue) ? defaultValue : runtimeValue; + } + } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/image/observation/DefaultImageModelObservationConventionTests.java b/spring-ai-core/src/test/java/org/springframework/ai/image/observation/DefaultImageModelObservationConventionTests.java index 6981e8fb2..d2ae8466e 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/image/observation/DefaultImageModelObservationConventionTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/image/observation/DefaultImageModelObservationConventionTests.java @@ -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), diff --git a/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelObservationContextTests.java b/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelObservationContextTests.java index 8ab69b6ef..febdc9435 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelObservationContextTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelObservationContextTests.java @@ -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(); diff --git a/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelPromptContentObservationFilterTests.java b/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelPromptContentObservationFilterTests.java index 1b6b0c522..122f97837 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelPromptContentObservationFilterTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelPromptContentObservationFilterTests.java @@ -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); diff --git a/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelRequestOptionsTests.java b/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelRequestOptionsTests.java deleted file mode 100644 index 59c444c38..000000000 --- a/spring-ai-core/src/test/java/org/springframework/ai/image/observation/ImageModelRequestOptionsTests.java +++ /dev/null @@ -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"); - } - -}