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:
committed by
Christian Tzolov
parent
80007d4d6c
commit
17ba1fc3ba
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -208,6 +208,7 @@ public class OpenAiImageOptions implements ImageOptions {
|
||||
this.size = this.width + "x" + this.height;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getStyle() {
|
||||
return this.style;
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -451,6 +451,11 @@ public class StabilityAiImageOptions implements ImageOptions {
|
||||
this.steps = steps;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getStyle() {
|
||||
return getStylePreset();
|
||||
}
|
||||
|
||||
public String getStylePreset() {
|
||||
return stylePreset;
|
||||
}
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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();
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
|
||||
}
|
||||
Reference in New Issue
Block a user