diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java index 35227909b..feeda8df4 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java @@ -32,7 +32,7 @@ import org.springframework.ai.anthropic.api.AnthropicApi.ChatCompletionResponse; import org.springframework.ai.anthropic.api.AnthropicApi.ContentBlock; import org.springframework.ai.anthropic.api.AnthropicApi.ContentBlock.ContentBlockType; import org.springframework.ai.anthropic.api.AnthropicApi.Role; -import org.springframework.ai.anthropic.metadata.AnthropicChatResponseMetadata; +import org.springframework.ai.anthropic.metadata.AnthropicChatResponseMetadataUtils; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; @@ -228,7 +228,7 @@ public class AnthropicChatModel extends AbstractToolCallSupport implements ChatResponseMetadata { - - protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, model: %3$s, usage: %4$s, rateLimit: %5$s }"; - - public static AnthropicChatResponseMetadata from(AnthropicApi.ChatCompletionResponse result) { - Assert.notNull(result, "Anthropic ChatCompletionResult must not be null"); - AnthropicUsage usage = AnthropicUsage.from(result.usage()); - return new AnthropicChatResponseMetadata(result.id(), result.model(), usage); - } - - private final String id; - - private final String model; - - @Nullable - private RateLimit rateLimit; - - private final Usage usage; - - protected AnthropicChatResponseMetadata(String id, String model, AnthropicUsage usage) { - this(id, model, usage, null); - } - - protected AnthropicChatResponseMetadata(String id, String model, AnthropicUsage usage, - @Nullable AnthropicRateLimit rateLimit) { - this.id = id; - this.model = model; - this.usage = usage; - this.rateLimit = rateLimit; - } - - @Override - public String getId() { - return this.id; - } - - @Override - public String getModel() { - return this.model; - } - - @Override - @Nullable - public RateLimit getRateLimit() { - RateLimit rl = this.rateLimit; - return rl != null ? rl : new EmptyRateLimit(); - } - - @Override - public Usage getUsage() { - Usage usage = this.usage; - return usage != null ? usage : new EmptyUsage(); - } - - public AnthropicChatResponseMetadata withRateLimit(RateLimit rateLimit) { - this.rateLimit = rateLimit; - return this; - } - - @Override - public String toString() { - return AI_METADATA_STRING.formatted(getClass().getName(), getId(), getModel(), getUsage(), getRateLimit()); - } - -} diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/metadata/AnthropicChatResponseMetadataUtils.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/metadata/AnthropicChatResponseMetadataUtils.java new file mode 100644 index 000000000..82ba3074b --- /dev/null +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/metadata/AnthropicChatResponseMetadataUtils.java @@ -0,0 +1,54 @@ +/* + * Copyright 2023 - 2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.ai.anthropic.metadata; + +import org.springframework.ai.anthropic.api.AnthropicApi; +import org.springframework.ai.chat.metadata.ChatResponseMetadata; +import org.springframework.ai.chat.metadata.EmptyRateLimit; +import org.springframework.ai.chat.metadata.EmptyUsage; +import org.springframework.ai.chat.metadata.RateLimit; +import org.springframework.ai.chat.metadata.Usage; +import org.springframework.lang.Nullable; +import org.springframework.util.Assert; + +import java.util.HashMap; + +/** + * {@link ChatResponseMetadata} implementation for {@literal AnthropicApi}. + * + * @author Christian Tzolov + * @author Thomas Vitale + * @see ChatResponseMetadata + * @see RateLimit + * @see Usage + * @since 1.0.0 + */ +public abstract class AnthropicChatResponseMetadataUtils { + + public static ChatResponseMetadata from(AnthropicApi.ChatCompletionResponse result) { + Assert.notNull(result, "Anthropic ChatCompletionResult must not be null"); + AnthropicUsage usage = AnthropicUsage.from(result.usage()); + return ChatResponseMetadata.builder() + .withId(result.id()) + .withModel(result.model()) + .withUsage(usage) + .withKeyValue("stop-reason", result.stopReason()) + .withKeyValue("stop-sequence", result.stopSequence()) + .withKeyValue("type", result.type()) + .build(); + } + +} diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java index b88345848..f28c320ab 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java @@ -43,7 +43,7 @@ import com.azure.core.util.BinaryData; import com.azure.core.util.IterableStream; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.azure.openai.metadata.AzureOpenAiChatResponseMetadata; +import org.springframework.ai.azure.openai.metadata.AzureOpenAiChatResponseMetadataUtils; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.metadata.PromptMetadata; @@ -155,7 +155,7 @@ public class AzureOpenAiChatModel extends PromptMetadata promptFilterMetadata = generatePromptMetadata(chatCompletions); return new ChatResponse(generations, - AzureOpenAiChatResponseMetadata.from(chatCompletions, promptFilterMetadata)); + AzureOpenAiChatResponseMetadataUtils.from(chatCompletions, promptFilterMetadata)); } @Override diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatResponseMetadata.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatResponseMetadataUtils.java similarity index 54% rename from models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatResponseMetadata.java rename to models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatResponseMetadataUtils.java index a20c40ae0..700066b69 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatResponseMetadata.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatResponseMetadataUtils.java @@ -22,8 +22,6 @@ import org.springframework.ai.chat.metadata.PromptMetadata; import org.springframework.ai.chat.metadata.Usage; import org.springframework.util.Assert; -import java.util.HashMap; - /** * {@link ChatResponseMetadata} implementation for * {@literal Microsoft Azure OpenAI Service}. @@ -33,51 +31,22 @@ import java.util.HashMap; * @see ChatResponseMetadata * @since 0.7.1 */ -public class AzureOpenAiChatResponseMetadata extends HashMap implements ChatResponseMetadata { - - protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, usage: %3$s, rateLimit: %4$s }"; +public abstract class AzureOpenAiChatResponseMetadataUtils { @SuppressWarnings("all") - public static AzureOpenAiChatResponseMetadata from(ChatCompletions chatCompletions, - PromptMetadata promptFilterMetadata) { + public static ChatResponseMetadata from(ChatCompletions chatCompletions, PromptMetadata promptFilterMetadata) { Assert.notNull(chatCompletions, "Azure OpenAI ChatCompletions must not be null"); String id = chatCompletions.getId(); AzureOpenAiUsage usage = AzureOpenAiUsage.from(chatCompletions); - AzureOpenAiChatResponseMetadata chatResponseMetadata = new AzureOpenAiChatResponseMetadata(id, usage, - promptFilterMetadata); + ChatResponseMetadata chatResponseMetadata = ChatResponseMetadata.builder() + .withId(id) + .withUsage(usage) + .withModel(chatCompletions.getModel()) + .withPromptMetadata(promptFilterMetadata) + .withKeyValue("system-fingerprint", chatCompletions.getSystemFingerprint()) + .build(); + return chatResponseMetadata; } - private final String id; - - private final Usage usage; - - private final PromptMetadata promptMetadata; - - protected AzureOpenAiChatResponseMetadata(String id, AzureOpenAiUsage usage, PromptMetadata promptMetadata) { - this.id = id; - this.usage = usage; - this.promptMetadata = promptMetadata; - } - - @Override - public String getId() { - return this.id; - } - - @Override - public Usage getUsage() { - return this.usage; - } - - @Override - public PromptMetadata getPromptMetadata() { - return this.promptMetadata; - } - - @Override - public String toString() { - return AI_METADATA_STRING.formatted(getClass().getTypeName(), getId(), getUsage(), getRateLimit()); - } - } diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiImageResponseMetadata.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiImageResponseMetadata.java index e821913f7..6d01d5cbb 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiImageResponseMetadata.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiImageResponseMetadata.java @@ -2,6 +2,7 @@ package org.springframework.ai.azure.openai.metadata; import com.azure.ai.openai.models.ImageGenerations; import org.springframework.ai.image.ImageResponseMetadata; +import org.springframework.ai.model.MutableResponseMetadata; import org.springframework.util.Assert; import java.util.HashMap; @@ -15,7 +16,7 @@ import java.util.Objects; * @author Benoit Moussaud * @since 1.0.0 M1 */ -public class AzureOpenAiImageResponseMetadata extends HashMap implements ImageResponseMetadata { +public class AzureOpenAiImageResponseMetadata extends ImageResponseMetadata { private final Long created; diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java index 4223b93d3..5a9dbda8e 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java @@ -30,7 +30,7 @@ import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionChunk; import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionMessage; import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionMessage.ToolCall; import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionRequest; -import org.springframework.ai.mistralai.metadata.MistralAiChatResponseMetadata; +import org.springframework.ai.mistralai.metadata.MistralAiChatResponseMetadataUtils; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.AbstractFunctionCallSupport; import org.springframework.ai.model.function.FunctionCallbackContext; @@ -120,7 +120,7 @@ public class MistralAiChatModel extends .withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), null))) .toList(); - return new ChatResponse(generations, MistralAiChatResponseMetadata.from(chatCompletion)); + return new ChatResponse(generations, MistralAiChatResponseMetadataUtils.from(chatCompletion)); }); } diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/metadata/MistralAiChatResponseMetadata.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/metadata/MistralAiChatResponseMetadata.java deleted file mode 100644 index 2ca8bfbf6..000000000 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/metadata/MistralAiChatResponseMetadata.java +++ /dev/null @@ -1,62 +0,0 @@ -package org.springframework.ai.mistralai.metadata; - -import org.springframework.ai.chat.metadata.ChatResponseMetadata; -import org.springframework.ai.chat.metadata.EmptyUsage; -import org.springframework.ai.chat.metadata.Usage; -import org.springframework.ai.mistralai.api.MistralAiApi; -import org.springframework.util.Assert; - -import java.util.HashMap; - -/** - * {@link ChatResponseMetadata} implementation for {@literal Mistral AI}. - * - * @author Thomas Vitale - * @see ChatResponseMetadata - * @see Usage - * @since 1.0.0 - */ -public class MistralAiChatResponseMetadata extends HashMap implements ChatResponseMetadata { - - protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, model: %3$s, usage: %4$s }"; - - public static MistralAiChatResponseMetadata from(MistralAiApi.ChatCompletion result) { - Assert.notNull(result, "Mistral AI ChatCompletion must not be null"); - MistralAiUsage usage = MistralAiUsage.from(result.usage()); - return new MistralAiChatResponseMetadata(result.id(), result.model(), usage); - } - - private final String id; - - private final String model; - - private final Usage usage; - - protected MistralAiChatResponseMetadata(String id, String model, MistralAiUsage usage) { - this.id = id; - this.model = model; - this.usage = usage; - } - - @Override - public String getId() { - return this.id; - } - - @Override - public String getModel() { - return this.model; - } - - @Override - public Usage getUsage() { - Usage usage = this.usage; - return usage != null ? usage : new EmptyUsage(); - } - - @Override - public String toString() { - return AI_METADATA_STRING.formatted(getClass().getName(), getId(), getModel(), getUsage()); - } - -} diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/metadata/MistralAiChatResponseMetadataUtils.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/metadata/MistralAiChatResponseMetadataUtils.java new file mode 100644 index 000000000..e86bbbcf8 --- /dev/null +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/metadata/MistralAiChatResponseMetadataUtils.java @@ -0,0 +1,32 @@ +package org.springframework.ai.mistralai.metadata; + +import org.springframework.ai.chat.metadata.ChatResponseMetadata; +import org.springframework.ai.chat.metadata.EmptyUsage; +import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.mistralai.api.MistralAiApi; +import org.springframework.util.Assert; + +import java.util.HashMap; + +/** + * {@link ChatResponseMetadata} implementation for {@literal Mistral AI}. + * + * @author Thomas Vitale + * @see ChatResponseMetadata + * @see Usage + * @since 1.0.0 + */ +public abstract class MistralAiChatResponseMetadataUtils { + + public static ChatResponseMetadata from(MistralAiApi.ChatCompletion result) { + Assert.notNull(result, "Mistral AI ChatCompletion must not be null"); + MistralAiUsage usage = MistralAiUsage.from(result.usage()); + return ChatResponseMetadata.builder() + .withId(result.id()) + .withModel(result.model()) + .withUsage(usage) + .withKeyValue("created", result.created()) + .build(); + } + +} diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java index 39405c2d6..886dced25 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java @@ -19,7 +19,7 @@ import java.util.Base64; import java.util.List; import org.springframework.ai.chat.model.ChatModel; -import org.springframework.ai.ollama.metadata.OllamaChatResponseMetadata; +import org.springframework.ai.ollama.metadata.OllamaChatResponseMetadataUtils; import reactor.core.publisher.Flux; import org.springframework.ai.chat.model.ChatResponse; @@ -102,7 +102,7 @@ public class OllamaChatModel implements ChatModel { if (response.promptEvalCount() != null && response.evalCount() != null) { generator = generator.withGenerationMetadata(ChatGenerationMetadata.from("unknown", null)); } - return new ChatResponse(List.of(generator), OllamaChatResponseMetadata.from(response)); + return new ChatResponse(List.of(generator), OllamaChatResponseMetadataUtils.from(response)); } @Override @@ -116,7 +116,7 @@ public class OllamaChatModel implements ChatModel { if (Boolean.TRUE.equals(chunk.done())) { generation = generation.withGenerationMetadata(ChatGenerationMetadata.from("unknown", null)); } - return new ChatResponse(List.of(generation), OllamaChatResponseMetadata.from(chunk)); + return new ChatResponse(List.of(generation), OllamaChatResponseMetadataUtils.from(chunk)); }); } diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatResponseMetadata.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatResponseMetadataUtils.java similarity index 61% rename from models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatResponseMetadata.java rename to models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatResponseMetadataUtils.java index 6f1d213fa..c98b197ad 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatResponseMetadata.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/metadata/OllamaChatResponseMetadataUtils.java @@ -28,30 +28,23 @@ import java.util.HashMap; * @see ChatResponseMetadata * @author Fu Cheng */ -public class OllamaChatResponseMetadata extends HashMap implements ChatResponseMetadata { +public abstract class OllamaChatResponseMetadataUtils { - protected static final String AI_METADATA_STRING = "{ @type: %1$s, usage: %2$s, rateLimit: %3$s }"; - - public static OllamaChatResponseMetadata from(OllamaApi.ChatResponse response) { + public static ChatResponseMetadata from(OllamaApi.ChatResponse response) { Assert.notNull(response, "OllamaApi.ChatResponse must not be null"); - Usage usage = OllamaUsage.from(response); - return new OllamaChatResponseMetadata(usage); - } + return ChatResponseMetadata.builder() + .withUsage(OllamaUsage.from(response)) + .withModel(response.model()) + .withKeyValue("created-at", response.createdAt()) + .withKeyValue("eval-duration", response.evalDuration()) + .withKeyValue("eval-count", response.evalCount()) + .withKeyValue("load-duration", response.loadDuration()) + .withKeyValue("eval-duration", response.promptEvalDuration()) + .withKeyValue("eval-count", response.promptEvalCount()) + .withKeyValue("total-duration", response.totalDuration()) + .withKeyValue("done", response.done()) + .build(); - private final Usage usage; - - protected OllamaChatResponseMetadata(Usage usage) { - this.usage = usage; - } - - @Override - public Usage getUsage() { - return this.usage; - } - - @Override - public String toString() { - return AI_METADATA_STRING.formatted(getClass().getTypeName(), getUsage(), getRateLimit()); } } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/ImageResponseMetadata.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/ImageResponseMetadata.java index f339b4a0b..3ec1ad510 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/ImageResponseMetadata.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/ImageResponseMetadata.java @@ -1,5 +1,5 @@ package org.springframework.ai.openai; -public class ImageResponseMetadata { +public interface ImageResponseMetadata { } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java index fca2de714..704cbd232 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java @@ -15,14 +15,6 @@ */ package org.springframework.ai.openai; -import java.util.ArrayList; -import java.util.Base64; -import java.util.HashSet; -import java.util.List; -import java.util.Map; -import java.util.Set; -import java.util.concurrent.ConcurrentHashMap; - import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.chat.messages.AssistantMessage; @@ -49,7 +41,7 @@ import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.ChatCom import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.MediaContent; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.ToolCall; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest; -import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadata; +import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadataUtils; import org.springframework.ai.openai.metadata.support.OpenAiResponseHeaderExtractor; import org.springframework.ai.retry.RetryUtils; import org.springframework.http.ResponseEntity; @@ -57,10 +49,17 @@ import org.springframework.retry.support.RetryTemplate; import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; import org.springframework.util.MimeType; - import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import java.util.ArrayList; +import java.util.Base64; +import java.util.HashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.concurrent.ConcurrentHashMap; + /** * {@link ChatModel} and {@link StreamingChatModel} implementation for {@literal OpenAI} * backed by {@link OpenAiApi}. @@ -165,7 +164,7 @@ public class OpenAiChatModel extends AbstractToolCallSupport imp } // Non function calling. - RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity); + RateLimit rateLimit = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity); List choices = chatCompletion.choices(); if (choices == null) { @@ -187,7 +186,7 @@ public class OpenAiChatModel extends AbstractToolCallSupport imp }).toList(); return new ChatResponse(generations, - OpenAiChatResponseMetadata.from(completionEntity.getBody()).withRateLimit(rateLimits)); + OpenAiChatResponseMetadataUtils.from(completionEntity.getBody(), rateLimit)); }); } @@ -237,7 +236,7 @@ public class OpenAiChatModel extends AbstractToolCallSupport imp }).toList(); if (chatCompletion2.usage() != null) { - return new ChatResponse(generations, OpenAiChatResponseMetadata.from(chatCompletion2)); + return new ChatResponse(generations, OpenAiChatResponseMetadataUtils.from(chatCompletion2)); } else { return new ChatResponse(generations); 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 efc4b24ec..7c4267fbe 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 @@ -15,8 +15,6 @@ */ package org.springframework.ai.openai; -import java.util.List; - import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.image.Image; @@ -29,12 +27,13 @@ import org.springframework.ai.image.ImageResponseMetadata; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.openai.api.OpenAiImageApi; import org.springframework.ai.openai.metadata.OpenAiImageGenerationMetadata; -import org.springframework.ai.openai.metadata.OpenAiImageResponseMetadata; import org.springframework.ai.retry.RetryUtils; import org.springframework.http.ResponseEntity; import org.springframework.retry.support.RetryTemplate; import org.springframework.util.Assert; +import java.util.List; + /** * OpenAiImageModel is a class that implements the ImageModel interface. It provides a * client for calling the OpenAI image generation API. @@ -130,7 +129,7 @@ public class OpenAiImageModel implements ImageModel { new OpenAiImageGenerationMetadata(entry.revisedPrompt())); }).toList(); - ImageResponseMetadata openAiImageResponseMetadata = OpenAiImageResponseMetadata.from(imageApiResponse); + ImageResponseMetadata openAiImageResponseMetadata = new ImageResponseMetadata(imageApiResponse.created()); return new ImageResponse(imageGenerationList, openAiImageResponseMetadata); } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadata.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadata.java deleted file mode 100644 index 93de30712..000000000 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadata.java +++ /dev/null @@ -1,103 +0,0 @@ -/* - * Copyright 2023 - 2024 the original author or authors. - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * https://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ -package org.springframework.ai.openai.metadata; - -import org.springframework.ai.chat.metadata.ChatResponseMetadata; -import org.springframework.ai.chat.metadata.EmptyRateLimit; -import org.springframework.ai.chat.metadata.EmptyUsage; -import org.springframework.ai.chat.metadata.RateLimit; -import org.springframework.ai.chat.metadata.Usage; -import org.springframework.ai.openai.api.OpenAiApi; -import org.springframework.lang.Nullable; -import org.springframework.util.Assert; - -import java.util.HashMap; - -/** - * {@link ChatResponseMetadata} implementation for {@literal OpenAI}. - * - * @author John Blum - * @author Thomas Vitale - * @see ChatResponseMetadata - * @see RateLimit - * @see Usage - * @since 0.7.0 - */ -public class OpenAiChatResponseMetadata extends HashMap implements ChatResponseMetadata { - - protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, model: %3$s, usage: %4$s, rateLimit: %5$s }"; - - public static OpenAiChatResponseMetadata from(OpenAiApi.ChatCompletion result) { - Assert.notNull(result, "OpenAI ChatCompletionResult must not be null"); - OpenAiUsage usage = OpenAiUsage.from(result.usage()); - return new OpenAiChatResponseMetadata(result.id(), result.model(), usage); - } - - private final String id; - - private final String model; - - @Nullable - private RateLimit rateLimit; - - private final Usage usage; - - protected OpenAiChatResponseMetadata(String id, String model, OpenAiUsage usage) { - this(id, model, usage, null); - } - - protected OpenAiChatResponseMetadata(String id, String model, OpenAiUsage usage, - @Nullable OpenAiRateLimit rateLimit) { - this.id = id; - this.model = model; - this.usage = usage; - this.rateLimit = rateLimit; - } - - @Override - public String getId() { - return this.id; - } - - @Override - public String getModel() { - return this.model; - } - - @Override - @Nullable - public RateLimit getRateLimit() { - RateLimit rateLimit = this.rateLimit; - return rateLimit != null ? rateLimit : new EmptyRateLimit(); - } - - @Override - public Usage getUsage() { - Usage usage = this.usage; - return usage != null ? usage : new EmptyUsage(); - } - - public OpenAiChatResponseMetadata withRateLimit(RateLimit rateLimit) { - this.rateLimit = rateLimit; - return this; - } - - @Override - public String toString() { - return AI_METADATA_STRING.formatted(getClass().getName(), getId(), getModel(), getUsage(), getRateLimit()); - } - -} diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadataUtils.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadataUtils.java new file mode 100644 index 000000000..ede428834 --- /dev/null +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadataUtils.java @@ -0,0 +1,59 @@ +/* + * Copyright 2023 - 2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.ai.openai.metadata; + +import org.springframework.ai.chat.metadata.ChatResponseMetadata; +import org.springframework.ai.chat.metadata.RateLimit; +import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.openai.api.OpenAiApi; +import org.springframework.util.Assert; + +/** + * {@link ChatResponseMetadata} implementation for {@literal OpenAI}. + * + * @author John Blum + * @author Thomas Vitale + * @see ChatResponseMetadata + * @see RateLimit + * @see Usage + * @since 0.7.0 + */ +public abstract class OpenAiChatResponseMetadataUtils { + + public static ChatResponseMetadata from(OpenAiApi.ChatCompletion result) { + Assert.notNull(result, "OpenAI ChatCompletionResult must not be null"); + return ChatResponseMetadata.builder() + .withId(result.id()) + .withUsage(OpenAiUsage.from(result.usage())) + .withModel(result.model()) + .withKeyValue("created", result.created()) + .withKeyValue("system-fingerprint", result.systemFingerprint()) + .build(); + } + + public static ChatResponseMetadata from(OpenAiApi.ChatCompletion result, RateLimit rateLimit) { + Assert.notNull(result, "OpenAI ChatCompletionResult must not be null"); + return ChatResponseMetadata.builder() + .withId(result.id()) + .withUsage(OpenAiUsage.from(result.usage())) + .withModel(result.model()) + .withRateLimit(rateLimit) + .withKeyValue("created", result.created()) + .withKeyValue("system-fingerprint", result.systemFingerprint()) + .build(); + } + +} diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiImageResponseMetadata.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiImageResponseMetadata.java deleted file mode 100644 index ec9519c82..000000000 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiImageResponseMetadata.java +++ /dev/null @@ -1,62 +0,0 @@ -/* - * Copyright 2023 - 2024 the original author or authors. - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * https://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ -package org.springframework.ai.openai.metadata; - -import org.springframework.ai.image.ImageResponseMetadata; -import org.springframework.ai.openai.api.OpenAiImageApi; -import org.springframework.util.Assert; - -import java.util.HashMap; -import java.util.Objects; - -public class OpenAiImageResponseMetadata extends HashMap implements ImageResponseMetadata { - - private final Long created; - - public static OpenAiImageResponseMetadata from(OpenAiImageApi.OpenAiImageResponse openAiImageResponse) { - Assert.notNull(openAiImageResponse, "OpenAiImageResponse must not be null"); - return new OpenAiImageResponseMetadata(openAiImageResponse.created()); - } - - protected OpenAiImageResponseMetadata(Long created) { - this.created = created; - } - - @Override - public Long getCreated() { - return this.created; - } - - @Override - public String toString() { - return "OpenAiImageResponseMetadata{" + "created=" + created + '}'; - } - - @Override - public boolean equals(Object o) { - if (this == o) - return true; - if (!(o instanceof OpenAiImageResponseMetadata that)) - return false; - return Objects.equals(created, that.created); - } - - @Override - public int hashCode() { - return Objects.hash(created); - } - -} diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/audio/OpenAiAudioSpeechResponseMetadata.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/audio/OpenAiAudioSpeechResponseMetadata.java index 4f38f3c0d..efcb6ebca 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/audio/OpenAiAudioSpeechResponseMetadata.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/audio/OpenAiAudioSpeechResponseMetadata.java @@ -18,6 +18,7 @@ package org.springframework.ai.openai.metadata.audio; import org.springframework.ai.chat.metadata.EmptyRateLimit; import org.springframework.ai.chat.metadata.RateLimit; +import org.springframework.ai.model.MutableResponseMetadata; import org.springframework.ai.model.ResponseMetadata; import org.springframework.ai.openai.api.OpenAiAudioApi; import org.springframework.lang.Nullable; @@ -31,7 +32,7 @@ import java.util.HashMap; * @author Ahmed Yousri * @see RateLimit */ -public class OpenAiAudioSpeechResponseMetadata extends HashMap implements ResponseMetadata { +public class OpenAiAudioSpeechResponseMetadata extends MutableResponseMetadata { protected static final String AI_METADATA_STRING = "{ @type: %1$s, requestsLimit: %2$s }"; diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/audio/OpenAiAudioTranscriptionResponseMetadata.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/audio/OpenAiAudioTranscriptionResponseMetadata.java index 5add8aa5b..5c20f831c 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/audio/OpenAiAudioTranscriptionResponseMetadata.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/audio/OpenAiAudioTranscriptionResponseMetadata.java @@ -17,6 +17,7 @@ package org.springframework.ai.openai.metadata.audio; import org.springframework.ai.chat.metadata.EmptyRateLimit; import org.springframework.ai.chat.metadata.RateLimit; +import org.springframework.ai.model.MutableResponseMetadata; import org.springframework.ai.model.ResponseMetadata; import org.springframework.ai.openai.api.OpenAiAudioApi; import org.springframework.ai.openai.metadata.OpenAiRateLimit; @@ -32,7 +33,7 @@ import java.util.HashMap; * @since 0.8.1 * @see RateLimit */ -public class OpenAiAudioTranscriptionResponseMetadata extends HashMap implements ResponseMetadata { +public class OpenAiAudioTranscriptionResponseMetadata extends MutableResponseMetadata { protected static final String AI_METADATA_STRING = "{ @type: %1$s, rateLimit: %4$s }"; diff --git a/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModel.java b/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModel.java index 2783f6df3..91e6f1292 100644 --- a/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModel.java +++ b/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModel.java @@ -200,9 +200,13 @@ public class PostgresMlEmbeddingModel extends AbstractEmbeddingModel implements } } - var metadata = new EmbeddingResponseMetadata( - Map.of("transformer", optionsToUse.getTransformer(), "vector-type", optionsToUse.getVectorType().name(), - "kwargs", ModelOptionsUtils.toJsonString(optionsToUse.getKwargs()))); + var metadata = new EmbeddingResponseMetadata(); + Map embeddingMetadata = Map.of("transformer", optionsToUse.getTransformer(), "vector-type", + optionsToUse.getVectorType().name(), "kwargs", + ModelOptionsUtils.toJsonString(optionsToUse.getKwargs())); + for (Map.Entry entry : embeddingMetadata.entrySet()) { + metadata.put(entry.getKey(), entry.getValue()); + } return new EmbeddingResponse(data, metadata); } diff --git a/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModelIT.java b/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModelIT.java index 55c05d06d..a6ec7e26b 100644 --- a/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModelIT.java +++ b/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModelIT.java @@ -30,6 +30,7 @@ import org.junit.jupiter.params.provider.ValueSource; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; +import org.springframework.ai.embedding.EmbeddingResponseMetadata; import org.springframework.ai.postgresml.PostgresMlEmbeddingModel.VectorType; import org.testcontainers.containers.PostgreSQLContainer; @@ -144,7 +145,21 @@ class PostgresMlEmbeddingModelIT { assertThat(embeddingResponse).isNotNull(); assertThat(embeddingResponse.getResults()).hasSize(3); - assertThat(embeddingResponse.getMetadata()).containsExactlyInAnyOrderEntriesOf( + EmbeddingResponseMetadata metadata = embeddingResponse.getMetadata(); + + assertThat(metadata.get("transformer").toString()) + .as("Transformer in metadata should be 'distilbert-base-uncased'") + .isEqualTo("distilbert-base-uncased"); + + assertThat(metadata.get("vector-type").toString()) + .as("Vector type in metadata should match expected vector type") + .isEqualTo(vectorType); + + assertThat(metadata.get("kwargs").toString()).as("kwargs in metadata should be '{}'").isEqualTo("{}"); + + assertThat(metadata.getRawMap().keySet()).as("Metadata should contain exactly the expected keys") + .containsExactlyInAnyOrder("transformer", "vector-type", "kwargs"); + assertThat(embeddingResponse.getMetadata().getRawMap()).containsExactlyInAnyOrderEntriesOf( Map.of("transformer", "distilbert-base-uncased", "vector-type", vectorType, "kwargs", "{}")); assertThat(embeddingResponse.getResults().get(0).getIndex()).isEqualTo(0); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(768); @@ -170,7 +185,7 @@ class PostgresMlEmbeddingModelIT { assertThat(embeddingResponse).isNotNull(); assertThat(embeddingResponse.getResults()).hasSize(3); - assertThat(embeddingResponse.getMetadata()).containsExactlyInAnyOrderEntriesOf(Map.of("transformer", + assertThat(embeddingResponse.getMetadata().getRawMap()).containsExactlyInAnyOrderEntriesOf(Map.of("transformer", "distilbert-base-uncased", "vector-type", VectorType.PG_VECTOR.name(), "kwargs", "{}")); assertThat(embeddingResponse.getResults().get(0).getIndex()).isEqualTo(0); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(768); @@ -192,7 +207,7 @@ class PostgresMlEmbeddingModelIT { assertThat(embeddingResponse).isNotNull(); assertThat(embeddingResponse.getResults()).hasSize(3); - assertThat(embeddingResponse.getMetadata()).containsExactlyInAnyOrderEntriesOf(Map.of("transformer", + assertThat(embeddingResponse.getMetadata().getRawMap()).containsExactlyInAnyOrderEntriesOf(Map.of("transformer", "intfloat/e5-small", "vector-type", VectorType.PG_ARRAY.name(), "kwargs", "{\"device\":\"cpu\"}")); assertThat(embeddingResponse.getResults().get(0).getIndex()).isEqualTo(0); 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 abb52c9a9..bf6f41f1a 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 @@ -119,7 +119,7 @@ public class StabilityAiImageModel implements ImageModel { new StabilityAiImageGenerationMetadata(entry.finishReason(), entry.seed())); }).toList(); - return new ImageResponse(imageGenerationList, ImageResponseMetadata.NULL); + return new ImageResponse(imageGenerationList, new ImageResponseMetadata()); } private StabilityAiImageOptions convertOptions(ImageOptions runtimeOptions) { diff --git a/models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingModelTests.java b/models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingModelTests.java index 40f963b5d..57ee908b3 100644 --- a/models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingModelTests.java +++ b/models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingModelTests.java @@ -24,6 +24,7 @@ import org.springframework.ai.document.Document; import org.springframework.ai.embedding.EmbeddingResponse; import static org.assertj.core.api.Assertions.assertThat; +import static org.junit.jupiter.api.Assertions.assertTrue; /** * @author Christian Tzolov @@ -76,7 +77,7 @@ public class TransformersEmbeddingModelTests { embeddingModel.afterPropertiesSet(); EmbeddingResponse embed = embeddingModel.embedForResponse(List.of("Hello world", "World is big")); assertThat(embed.getResults()).hasSize(2); - assertThat(embed.getMetadata()).isEmpty(); + assertTrue(embed.getMetadata().isEmpty(), "Expected embed metadata to be empty, but it was not."); assertThat(embed.getResults().get(0).getOutput()).hasSize(384); assertThat(DF.format(embed.getResults().get(0).getOutput().get(0))).isEqualTo(DF.format(-0.19744634628295898)); diff --git a/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/VertexAiEmbeddingUsage.java b/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/VertexAiEmbeddingUsage.java new file mode 100644 index 000000000..ef0152c23 --- /dev/null +++ b/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/VertexAiEmbeddingUsage.java @@ -0,0 +1,28 @@ +package org.springframework.ai.vertexai.embedding; + +import org.springframework.ai.chat.metadata.Usage; + +public class VertexAiEmbeddingUsage implements Usage { + + private final Integer totalTokens; + + public VertexAiEmbeddingUsage(Integer totalTokens) { + this.totalTokens = totalTokens; + } + + @Override + public Long getPromptTokens() { + return 0L; + } + + @Override + public Long getGenerationTokens() { + return 0L; + } + + @Override + public Long getTotalTokens() { + return Long.valueOf(this.totalTokens); + } + +} diff --git a/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/multimodal/VertexAiMultimodalEmbeddingModel.java b/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/multimodal/VertexAiMultimodalEmbeddingModel.java index e4c3132ff..8e55b347b 100644 --- a/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/multimodal/VertexAiMultimodalEmbeddingModel.java +++ b/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/multimodal/VertexAiMultimodalEmbeddingModel.java @@ -24,6 +24,7 @@ import com.google.protobuf.Value; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.chat.messages.Media; +import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.DocumentEmbeddingModel; import org.springframework.ai.embedding.DocumentEmbeddingRequest; @@ -35,6 +36,7 @@ import org.springframework.ai.embedding.EmbeddingResultMetadata; import org.springframework.ai.embedding.EmbeddingResultMetadata.ModalityType; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.vertexai.embedding.VertexAiEmbeddigConnectionDetails; +import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUsage; import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils; import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils.ImageBuilder; import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils.MultimodalInstanceBuilder; @@ -239,10 +241,11 @@ public class VertexAiMultimodalEmbeddingModel implements DocumentEmbeddingModel } - private EmbeddingResponseMetadata generateResponseMetadata(String model, Integer tokenCount) { + private EmbeddingResponseMetadata generateResponseMetadata(String model, Integer totalTokens) { EmbeddingResponseMetadata metadata = new EmbeddingResponseMetadata(); - metadata.put("model", model); - metadata.put("total-tokens", tokenCount); + metadata.setModel(model); + Usage usage = new VertexAiEmbeddingUsage(totalTokens); + metadata.setUsage(usage); return metadata; } diff --git a/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModel.java b/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModel.java index 70364a439..aba0dc2ee 100644 --- a/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModel.java +++ b/models/spring-ai-vertex-ai-embedding/src/main/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModel.java @@ -20,6 +20,7 @@ import com.google.cloud.aiplatform.v1.PredictRequest; import com.google.cloud.aiplatform.v1.PredictResponse; import com.google.cloud.aiplatform.v1.PredictionServiceClient; import com.google.protobuf.Value; +import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; @@ -32,6 +33,7 @@ import org.springframework.ai.vertexai.embedding.VertexAiEmbeddigConnectionDetai import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils; import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils.TextInstanceBuilder; import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils.TextParametersBuilder; +import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUsage; import org.springframework.util.Assert; import org.springframework.util.StringUtils; @@ -135,10 +137,11 @@ public class VertexAiTextEmbeddingModel extends AbstractEmbeddingModel { } } - private EmbeddingResponseMetadata generateResponseMetadata(String model, Integer tokenCount) { + private EmbeddingResponseMetadata generateResponseMetadata(String model, Integer totalTokens) { EmbeddingResponseMetadata metadata = new EmbeddingResponseMetadata(); - metadata.put("model", model); - metadata.put("total-tokens", tokenCount); + metadata.setModel(model); + Usage usage = new VertexAiEmbeddingUsage(totalTokens); + metadata.setUsage(usage); return metadata; } diff --git a/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/multimodal/VertexAiMultimodelEmbeddingModelIT.java b/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/multimodal/VertexAiMultimodelEmbeddingModelIT.java index 652f8abee..ca5177cd7 100644 --- a/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/multimodal/VertexAiMultimodelEmbeddingModelIT.java +++ b/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/multimodal/VertexAiMultimodelEmbeddingModelIT.java @@ -68,8 +68,13 @@ class VertexAiMultimodelEmbeddingModelIT { .isEqualTo(embeddingRequest.getInstructions().get(1).getId()); assertThat(embeddingResponse.getResults().get(1).getOutput()).hasSize(1408); - assertThat(embeddingResponse.getMetadata()).containsEntry("model", "multimodalembedding@001"); - assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 0); + assertThat(embeddingResponse.getMetadata().getModel()) + .as("Model in metadata should be 'multimodalembedding@001'") + .isEqualTo("multimodalembedding@001"); + + assertThat(embeddingResponse.getMetadata().getUsage().getTotalTokens()) + .as("Total tokens in metadata should be 0") + .isEqualTo("0"); assertThat(multiModelEmbeddingModel.dimensions()).isEqualTo(1408); } @@ -90,8 +95,8 @@ class VertexAiMultimodelEmbeddingModelIT { .isEqualTo(MimeTypeUtils.TEXT_PLAIN); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1408); - assertThat(embeddingResponse.getMetadata()).containsEntry("model", "multimodalembedding@001"); - assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 0); + assertThat(embeddingResponse.getMetadata().getModel()).isEqualTo("multimodalembedding@001"); + assertThat(embeddingResponse.getMetadata().getUsage()).isEqualTo(0); assertThat(multiModelEmbeddingModel.dimensions()).isEqualTo(1408); } @@ -113,8 +118,8 @@ class VertexAiMultimodelEmbeddingModelIT { .isEqualTo(MimeTypeUtils.TEXT_PLAIN); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1408); - assertThat(embeddingResponse.getMetadata()).containsEntry("model", "multimodalembedding@001"); - assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 0); + assertThat(embeddingResponse.getMetadata().getModel()).isEqualTo("multimodalembedding@001"); + assertThat(embeddingResponse.getMetadata().getUsage().getTotalTokens()).isEqualTo(0); assertThat(multiModelEmbeddingModel.dimensions()).isEqualTo(1408); } @@ -139,8 +144,8 @@ class VertexAiMultimodelEmbeddingModelIT { assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1408); - assertThat(embeddingResponse.getMetadata()).containsEntry("model", "multimodalembedding@001"); - assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 0); + assertThat(embeddingResponse.getMetadata().getModel()).isEqualTo("multimodalembedding@001"); + assertThat(embeddingResponse.getMetadata().getUsage().getTotalTokens()).isEqualTo(0); assertThat(multiModelEmbeddingModel.dimensions()).isEqualTo(1408); } @@ -164,8 +169,8 @@ class VertexAiMultimodelEmbeddingModelIT { .isEqualTo(new MimeType("video", "mp4")); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1408); - assertThat(embeddingResponse.getMetadata()).containsEntry("model", "multimodalembedding@001"); - assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 0); + assertThat(embeddingResponse.getMetadata().getModel()).isEqualTo("multimodalembedding@001"); + assertThat(embeddingResponse.getMetadata().getUsage().getTotalTokens()).isEqualTo(0); assertThat(multiModelEmbeddingModel.dimensions()).isEqualTo(1408); } @@ -198,8 +203,8 @@ class VertexAiMultimodelEmbeddingModelIT { .isEqualTo(EmbeddingResultMetadata.ModalityType.VIDEO); assertThat(embeddingResponse.getResults().get(2).getOutput()).hasSize(1408); - assertThat(embeddingResponse.getMetadata()).containsEntry("model", "multimodalembedding@001"); - assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 0); + assertThat(embeddingResponse.getMetadata().getModel()).isEqualTo("multimodalembedding@001"); + assertThat(embeddingResponse.getMetadata().getUsage()).isEqualTo(0); assertThat(multiModelEmbeddingModel.dimensions()).isEqualTo(1408); } diff --git a/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModelIT.java b/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModelIT.java index 43a6c86cc..4c9a9cdcb 100644 --- a/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModelIT.java +++ b/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingModelIT.java @@ -53,8 +53,12 @@ class VertexAiTextEmbeddingModelIT { assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(768); assertThat(embeddingResponse.getResults().get(1).getOutput()).hasSize(768); - assertThat(embeddingResponse.getMetadata()).containsEntry("model", modelName); - assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 5); + assertThat(embeddingResponse.getMetadata().getModel()).as("Model name in metadata should match expected model") + .isEqualTo(modelName); + + assertThat(embeddingResponse.getMetadata().getUsage().getTotalTokens()) + .as("Total tokens in metadata should be 5") + .isEqualTo(5L); assertThat(embeddingModel.dimensions()).isEqualTo(768); } diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java index 04e75de56..493941b8c 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java @@ -29,6 +29,7 @@ import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.ToolResponseMessage; import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.metadata.ChatResponseMetadata; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.Generation; @@ -38,7 +39,7 @@ import org.springframework.ai.model.ChatModelDescription; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.AbstractToolCallSupport; import org.springframework.ai.model.function.FunctionCallbackContext; -import org.springframework.ai.vertexai.gemini.metadata.VertexAiChatResponseMetadata; +import org.springframework.ai.vertexai.gemini.metadata.VertexAiChatResponseMetadataUtils; import org.springframework.ai.vertexai.gemini.metadata.VertexAiUsage; import org.springframework.beans.factory.DisposableBean; import org.springframework.lang.NonNull; @@ -244,8 +245,8 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements ChatResponseMetadata { +public abstract class VertexAiChatResponseMetadataUtils { - private final VertexAiUsage usage; - - public VertexAiChatResponseMetadata(VertexAiUsage usage) { - this.usage = usage; - } - - @Override - public Usage getUsage() { - return this.usage; + public static ChatResponseMetadata from(GenerateContentResponse.UsageMetadata usageMetadata) { + Assert.notNull(usageMetadata, "GenerateContentResponse.UsageMetadata must not be null"); + return ChatResponseMetadata.builder().withUsage(new VertexAiUsage(usageMetadata)).build(); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/metadata/ChatResponseMetadata.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/metadata/ChatResponseMetadata.java index 17fea86cd..4c4b54464 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/metadata/ChatResponseMetadata.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/metadata/ChatResponseMetadata.java @@ -15,40 +15,46 @@ */ package org.springframework.ai.chat.metadata; -import org.springframework.ai.model.ResponseMetadata; +import org.springframework.ai.model.MutableResponseMetadata; -import java.util.HashMap; +import java.util.Objects; /** - * Abstract Data Type (ADT) modeling common AI provider metadata returned in an AI - * response. + * Models common AI provider metadata returned in an AI response. * * @author John Blum * @author Thomas Vitale * @since 0.7.0 */ -public interface ChatResponseMetadata extends ResponseMetadata { +public class ChatResponseMetadata extends MutableResponseMetadata { - class DefaultChatResponseMetadata extends HashMap implements ChatResponseMetadata { + private static final String AI_METADATA_STRING = "{ id: %2$s, usage: %3$s, rateLimit: %4$s }"; - } + private String id = ""; // Set to blank to preserve backward compat with previous + // interface default methods - ChatResponseMetadata NULL = new DefaultChatResponseMetadata(); + private String model = ""; + + private RateLimit rateLimit = new EmptyRateLimit(); + + private Usage usage = new EmptyUsage(); + + private PromptMetadata promptMetadata = PromptMetadata.empty(); /** * A unique identifier for the chat completion operation. * @return unique operation identifier. */ - default String getId() { - return ""; + public String getId() { + return this.id; } /** * The model that handled the request. * @return the model that handled the request. */ - default String getModel() { - return ""; + public String getModel() { + return this.model; } /** @@ -56,8 +62,8 @@ public interface ChatResponseMetadata extends ResponseMetadata { * @return AI provider specific metadata on rate limits. * @see RateLimit */ - default RateLimit getRateLimit() { - return new EmptyRateLimit(); + public RateLimit getRateLimit() { + return this.rateLimit; } /** @@ -65,12 +71,85 @@ public interface ChatResponseMetadata extends ResponseMetadata { * @return AI provider specific metadata on API usage. * @see Usage */ - default Usage getUsage() { - return new EmptyUsage(); + public Usage getUsage() { + return this.usage; } - default PromptMetadata getPromptMetadata() { - return PromptMetadata.empty(); + /** + * Returns the prompt metadata gathered by the AI during request processing. + * @return the prompt metadata. + */ + public PromptMetadata getPromptMetadata() { + return this.promptMetadata; + } + + public static class Builder { + + private final ChatResponseMetadata chatResponseMetadata; + + public Builder() { + this.chatResponseMetadata = new ChatResponseMetadata(); + } + + public Builder withKeyValue(String key, Object value) { + this.chatResponseMetadata.put(key, value); + return this; + } + + public Builder withId(String id) { + this.chatResponseMetadata.id = id; + return this; + } + + public Builder withModel(String model) { + this.chatResponseMetadata.model = model; + return this; + } + + public Builder withRateLimit(RateLimit rateLimit) { + this.chatResponseMetadata.rateLimit = rateLimit; + return this; + } + + public Builder withUsage(Usage usage) { + this.chatResponseMetadata.usage = usage; + return this; + } + + public Builder withPromptMetadata(PromptMetadata promptMetadata) { + this.chatResponseMetadata.promptMetadata = promptMetadata; + return this; + } + + public ChatResponseMetadata build() { + return this.chatResponseMetadata; + } + + } + + public static Builder builder() { + return new Builder(); + } + + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (!(o instanceof ChatResponseMetadata that)) + return false; + return Objects.equals(id, that.id) && Objects.equals(model, that.model) + && Objects.equals(rateLimit, that.rateLimit) && Objects.equals(usage, that.usage) + && Objects.equals(promptMetadata, that.promptMetadata); + } + + @Override + public int hashCode() { + return Objects.hash(id, model, rateLimit, usage, promptMetadata); + } + + @Override + public String toString() { + return AI_METADATA_STRING.formatted(getId(), getUsage(), getRateLimit()); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ChatResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ChatResponse.java index 8c31776c0..0e06f9ecf 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ChatResponse.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ChatResponse.java @@ -40,7 +40,7 @@ public class ChatResponse implements ModelResponse { * provider. */ public ChatResponse(List generations) { - this(generations, ChatResponseMetadata.NULL); + this(generations, new ChatResponseMetadata()); } /** diff --git a/spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingResponseMetadata.java b/spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingResponseMetadata.java index 40a305838..26bb8bfd3 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingResponseMetadata.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingResponseMetadata.java @@ -21,6 +21,7 @@ import java.util.Map; import org.springframework.ai.chat.metadata.EmptyUsage; import org.springframework.ai.chat.metadata.Usage; +import org.springframework.ai.model.MutableResponseMetadata; import org.springframework.ai.model.ResponseMetadata; /** @@ -29,10 +30,7 @@ import org.springframework.ai.model.ResponseMetadata; * @author Christian Tzolov * @author Thomas Vitale */ -public class EmbeddingResponseMetadata extends HashMap implements ResponseMetadata { - - @Serial - private static final long serialVersionUID = 1L; +public class EmbeddingResponseMetadata extends MutableResponseMetadata { private String model; @@ -42,12 +40,15 @@ public class EmbeddingResponseMetadata extends HashMap implement } public EmbeddingResponseMetadata(String model, Usage usage) { - this.model = model; - this.usage = usage; + this(model, usage, Map.of()); } - public EmbeddingResponseMetadata(Map metadata) { - super(metadata); + public EmbeddingResponseMetadata(String model, Usage usage, Map metadata) { + this.model = model; + this.usage = usage; + for (Map.Entry entry : metadata.entrySet()) { + this.put(entry.getKey(), entry.getValue()); + } } /** diff --git a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageResponse.java index 70a4b946a..b6d6c87b8 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageResponse.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageResponse.java @@ -43,7 +43,7 @@ public class ImageResponse implements ModelResponse { * provider. */ public ImageResponse(List generations) { - this(generations, ImageResponseMetadata.NULL); + this(generations, new ImageResponseMetadata()); } /** diff --git a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageResponseMetadata.java b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageResponseMetadata.java index a80b31bba..fe5b78985 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageResponseMetadata.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageResponseMetadata.java @@ -15,6 +15,7 @@ */ package org.springframework.ai.image; +import org.springframework.ai.model.MutableResponseMetadata; import org.springframework.ai.model.ResponseMetadata; import java.util.HashMap; @@ -28,16 +29,20 @@ import java.util.HashMap; * @author Thomas Vitale * @since 1.0.0 */ -public interface ImageResponseMetadata extends ResponseMetadata { +public class ImageResponseMetadata extends MutableResponseMetadata { - class DefaultImageResponseMetadata extends HashMap implements ImageResponseMetadata { + private Long created; + public ImageResponseMetadata() { + this.created = System.currentTimeMillis(); } - ImageResponseMetadata NULL = new DefaultImageResponseMetadata(); + public ImageResponseMetadata(Long created) { + this.created = created; + } - default Long getCreated() { - return System.currentTimeMillis(); + public Long getCreated() { + return this.created; } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/MutableResponseMetadata.java b/spring-ai-core/src/main/java/org/springframework/ai/model/MutableResponseMetadata.java new file mode 100644 index 000000000..b16181a57 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/MutableResponseMetadata.java @@ -0,0 +1,115 @@ +package org.springframework.ai.model; + +import io.micrometer.common.lang.NonNull; +import io.micrometer.common.lang.Nullable; + +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import java.util.function.Function; + +public class MutableResponseMetadata implements ResponseMetadata { + + private final Map map = new ConcurrentHashMap<>(); + + /** + * Puts an element to the context. + * @param key key + * @param object value + * @param value type + * @return this for chaining + */ + public MutableResponseMetadata put(String key, T object) { + this.map.put(key, object); + return this; + } + + /** + * Gets an entry from the context. Returns {@code null} when entry is not present. + * @param key key + * @param value type + * @return entry or {@code null} if not present + */ + @Override + @Nullable + public T get(String key) { + return (T) this.map.get(key); + } + + /** + * Removes an entry from the context. + * @param key key by which to remove an entry + * @return the previous value associated with the key, or null if there was no mapping + * for the key + */ + public Object remove(Object key) { + return this.map.remove(key); + } + + /** + * Gets an entry from the context. Throws exception when entry is not present. + * @param key key + * @param value type + * @throws IllegalArgumentException if not present + * @return entry + */ + @Override + @NonNull + public T getRequired(Object key) { + T object = (T) this.map.get(key); + if (object == null) { + throw new IllegalArgumentException("Context does not have an entry for key [" + key + "]"); + } + return object; + } + + /** + * Checks if context contains a key. + * @param key key + * @return {@code true} when the context contains the entry with the given key + */ + @Override + public boolean containsKey(Object key) { + return this.map.containsKey(key); + } + + /** + * Returns an element or default if not present. + * @param key key + * @param defaultObject default object to return + * @param value type + * @return object or default if not present + */ + @Override + public T getOrDefault(Object key, T defaultObject) { + return (T) this.map.getOrDefault(key, defaultObject); + } + + @Override + public boolean isEmpty() { + return this.map.isEmpty(); + } + + /** + * Returns an element or calls a mapping function if entry not present. The function + * will insert the value to the map. + * @param key key + * @param mappingFunction mapping function + * @param value type + * @return object or one derived from the mapping function if not present + */ + public T computeIfAbsent(String key, Function mappingFunction) { + return (T) this.map.computeIfAbsent(key, mappingFunction); + } + + /** + * Clears the entries from the context. + */ + public void clear() { + this.map.clear(); + } + + public Map getRawMap() { + return map; + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/ResponseMetadata.java b/spring-ai-core/src/main/java/org/springframework/ai/model/ResponseMetadata.java index b1516ec3f..6309fa6f6 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/ResponseMetadata.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/ResponseMetadata.java @@ -15,7 +15,11 @@ */ package org.springframework.ai.model; -import java.util.Map; +import io.micrometer.common.KeyValues; +import io.micrometer.common.lang.NonNull; +import io.micrometer.common.lang.Nullable; + +import java.util.function.Supplier; /** * Interface representing metadata associated with an AI model's response. This interface @@ -27,6 +31,60 @@ import java.util.Map; * @author Mark Pollack * @since 0.8.0 */ -public interface ResponseMetadata extends Map { +public interface ResponseMetadata { + + /** + * Gets an entry from the context. Returns {@code null} when entry is not present. + * @param key key + * @param value type + * @return entry or {@code null} if not present + */ + @Nullable + T get(String key); + + /** + * Gets an entry from the context. Throws exception when entry is not present. + * @param key key + * @param value type + * @throws IllegalArgumentException if not present + * @return entry + */ + @NonNull + T getRequired(Object key); + + /** + * Checks if context contains a key. + * @param key key + * @return {@code true} when the context contains the entry with the given key + */ + boolean containsKey(Object key); + + /** + * Returns an element or default if not present. + * @param key key + * @param defaultObject default object to return + * @param value type + * @return object or default if not present + */ + T getOrDefault(Object key, T defaultObject); + + /** + * Returns an element or default if not present. + * @param key key + * @param defaultObjectSupplier supplier for default object to return + * @param value type + * @return object or default if not present + * @since 1.11.0 + */ + default T getOrDefault(String key, Supplier defaultObjectSupplier) { + T value = get(key); + return value != null ? value : defaultObjectSupplier.get(); + } + + /** + * Returns {@code true} if this map contains no key-value mappings. + * @return {@code true} if this map contains no key-value mappings + */ + boolean isEmpty(); } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/client/ChatClientResponseEntityTests.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/client/ChatClientResponseEntityTests.java index 2fae3d9f2..31fabb138 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/client/ChatClientResponseEntityTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/client/ChatClientResponseEntityTests.java @@ -29,7 +29,6 @@ import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.metadata.ChatResponseMetadata; -import org.springframework.ai.chat.metadata.ChatResponseMetadata.DefaultChatResponseMetadata; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.Generation; @@ -58,7 +57,7 @@ public class ChatClientResponseEntityTests { @Test public void responseEntityTest() { - ChatResponseMetadata metadata = new DefaultChatResponseMetadata(); + ChatResponseMetadata metadata = new ChatResponseMetadata(); metadata.put("key1", "value1"); var chatResponse = new ChatResponse(List.of(new Generation(""" @@ -75,7 +74,7 @@ public class ChatClientResponseEntityTests { .responseEntity(MyBean.class); assertThat(responseEntity.getResponse()).isEqualTo(chatResponse); - assertThat(responseEntity.getResponse().getMetadata().get("key1")).isEqualTo("value1"); + assertThat(responseEntity.getResponse().getMetadata().get("key1").toString()).isEqualTo("value1"); assertThat(responseEntity.getEntity()).isEqualTo(new MyBean("John", 30)); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/embedding/VertexAiTextEmbeddingModelAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/embedding/VertexAiTextEmbeddingModelAutoConfigurationIT.java index ca030ab9e..b4385d64f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/embedding/VertexAiTextEmbeddingModelAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/embedding/VertexAiTextEmbeddingModelAutoConfigurationIT.java @@ -112,8 +112,8 @@ public class VertexAiTextEmbeddingModelAutoConfigurationIT { .isEqualTo(EmbeddingResultMetadata.ModalityType.TEXT); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1408); - assertThat(embeddingResponse.getMetadata()).containsEntry("model", "multimodalembedding@001"); - assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 0); + assertThat(embeddingResponse.getMetadata().getModel()).isEqualTo("multimodalembedding@001"); + assertThat(embeddingResponse.getMetadata().getUsage()).isEqualTo(0); assertThat(multiModelEmbeddingModel.dimensions()).isEqualTo(1408);