Refactor Usage handling

- Remove model specific Usage implementations
- Add `Object getNativeUsage()` to Usage interface
  - This will allow the model specific Usage data to be returned
  - At the client side, client needs to cast the return type of `getNativeUsage` into the corresponding Usage returned by the model API
- Rename `generationTokens` to `completionTokens`
  - Since `completion` token name is more common among the models, renaming generation tokens into completion tokens
  - Maintain JSON deserialization compatibility for legacy `generationTokens` field
  - Remove deprecated Long-based constructors to avoid API ambiguity
- Change the prompt, completion and total token return types from Long to Integer
  - This is a breaking change that requires updating all constructor calls
  - Integer is sufficient for token counts and aligns better with most model APIs
- Use DefaultUsage for most of the model specific usage handling
  - When initializing set the native usage to the model specific usage type
  - Ensure immutability by making all fields final and removing setters
  - Add comprehensive test coverage for all functionality including edge cases

Resolves #1407
This commit is contained in:
Ilayaperumal Gopinathan
2025-01-23 14:28:02 +00:00
committed by Mark Pollack
parent 840304955e
commit 4b64aa0ca6
281 changed files with 825 additions and 1385 deletions

View File

@@ -40,13 +40,13 @@ import org.springframework.ai.anthropic.api.AnthropicApi.ContentBlock;
import org.springframework.ai.anthropic.api.AnthropicApi.ContentBlock.Source;
import org.springframework.ai.anthropic.api.AnthropicApi.ContentBlock.Type;
import org.springframework.ai.anthropic.api.AnthropicApi.Role;
import org.springframework.ai.anthropic.metadata.AnthropicUsage;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.MessageType;
import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.metadata.UsageUtils;
@@ -237,7 +237,7 @@ public class AnthropicChatModel extends AbstractToolCallSupport implements ChatM
AnthropicApi.ChatCompletionResponse completionResponse = completionEntity.getBody();
AnthropicApi.Usage usage = completionResponse.usage();
Usage currentChatResponseUsage = usage != null ? AnthropicUsage.from(completionResponse.usage())
Usage currentChatResponseUsage = usage != null ? this.getDefaultUsage(completionResponse.usage())
: new EmptyUsage();
Usage accumulatedUsage = UsageUtils.getCumulativeUsage(currentChatResponseUsage, previousChatResponse);
@@ -256,6 +256,11 @@ public class AnthropicChatModel extends AbstractToolCallSupport implements ChatM
return response;
}
private DefaultUsage getDefaultUsage(AnthropicApi.Usage usage) {
return new DefaultUsage(usage.inputTokens(), usage.outputTokens(), usage.inputTokens() + usage.outputTokens(),
usage);
}
@Override
public Flux<ChatResponse> stream(Prompt prompt) {
return this.internalStream(prompt, null);
@@ -282,7 +287,7 @@ public class AnthropicChatModel extends AbstractToolCallSupport implements ChatM
// @formatter:off
Flux<ChatResponse> chatResponseFlux = response.switchMap(chatCompletionResponse -> {
AnthropicApi.Usage usage = chatCompletionResponse.usage();
Usage currentChatResponseUsage = usage != null ? AnthropicUsage.from(chatCompletionResponse.usage()) : new EmptyUsage();
Usage currentChatResponseUsage = usage != null ? this.getDefaultUsage(chatCompletionResponse.usage()) : new EmptyUsage();
Usage accumulatedUsage = UsageUtils.getCumulativeUsage(currentChatResponseUsage, previousChatResponse);
ChatResponse chatResponse = toChatResponse(chatCompletionResponse, accumulatedUsage);
@@ -352,7 +357,7 @@ public class AnthropicChatModel extends AbstractToolCallSupport implements ChatM
}
private ChatResponseMetadata from(AnthropicApi.ChatCompletionResponse result) {
return from(result, AnthropicUsage.from(result.usage()));
return from(result, this.getDefaultUsage(result.usage()));
}
private ChatResponseMetadata from(AnthropicApi.ChatCompletionResponse result, Usage usage) {

View File

@@ -1,66 +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.anthropic.metadata;
import org.springframework.ai.anthropic.api.AnthropicApi;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal AnthropicApi}.
*
* @author Christian Tzolov
* @since 1.0.0
*/
public class AnthropicUsage implements Usage {
private final AnthropicApi.Usage usage;
protected AnthropicUsage(AnthropicApi.Usage usage) {
Assert.notNull(usage, "AnthropicApi Usage must not be null");
this.usage = usage;
}
public static AnthropicUsage from(AnthropicApi.Usage usage) {
return new AnthropicUsage(usage);
}
protected AnthropicApi.Usage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
return getUsage().inputTokens().longValue();
}
@Override
public Long getGenerationTokens() {
return getUsage().outputTokens().longValue();
}
@Override
public Long getTotalTokens() {
return this.getPromptTokens() + this.getGenerationTokens();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -84,7 +84,7 @@ class AnthropicChatModelIT {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(model);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}
@@ -100,11 +100,11 @@ class AnthropicChatModelIT {
AnthropicChatOptions.builder().model(modelName).build());
ChatResponse response = this.chatModel.call(prompt);
assertThat(response.getResults()).hasSize(1);
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isGreaterThan(0);
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isGreaterThan(0);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isGreaterThan(0);
assertThat(response.getMetadata().getUsage().getTotalTokens())
.isEqualTo(response.getMetadata().getUsage().getPromptTokens()
+ response.getMetadata().getUsage().getGenerationTokens());
+ response.getMetadata().getUsage().getCompletionTokens());
Generation generation = response.getResults().get(0);
assertThat(generation.getOutput().getText()).contains("Blackbeard");
assertThat(generation.getMetadata().getFinishReason()).isEqualTo("end_turn");
@@ -139,11 +139,11 @@ class AnthropicChatModelIT {
var referenceTokenUsage = this.chatModel.call(prompt).getMetadata().getUsage();
assertThat(streamingTokenUsage.getPromptTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getGenerationTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getCompletionTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isEqualTo(referenceTokenUsage.getPromptTokens());
// assertThat(streamingTokenUsage.getGenerationTokens()).isEqualTo(referenceTokenUsage.getGenerationTokens());
// assertThat(streamingTokenUsage.getCompletionTokens()).isEqualTo(referenceTokenUsage.getCompletionTokens());
// assertThat(streamingTokenUsage.getTotalTokens()).isEqualTo(referenceTokenUsage.getTotalTokens());
}

View File

@@ -147,7 +147,7 @@ public class AnthropicChatModelObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

View File

@@ -60,13 +60,13 @@ import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.contextpropagation.ObservationThreadLocalAccessor;
import reactor.core.publisher.Flux;
import org.springframework.ai.azure.openai.metadata.AzureOpenAiUsage;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.metadata.PromptMetadata;
import org.springframework.ai.chat.metadata.PromptMetadata.PromptFilterMetadata;
@@ -194,13 +194,14 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements Cha
}
public static ChatResponseMetadata from(ChatCompletions chatCompletions, PromptMetadata promptFilterMetadata) {
Usage usage = (chatCompletions.getUsage() != null) ? AzureOpenAiUsage.from(chatCompletions) : new EmptyUsage();
Usage usage = (chatCompletions.getUsage() != null) ? getDefaultUsage(chatCompletions.getUsage())
: new EmptyUsage();
return from(chatCompletions, promptFilterMetadata, usage);
}
public static ChatResponseMetadata from(ChatCompletions chatCompletions, PromptMetadata promptFilterMetadata,
CompletionsUsage usage) {
return from(chatCompletions, promptFilterMetadata, AzureOpenAiUsage.from(usage));
return from(chatCompletions, promptFilterMetadata, getDefaultUsage(usage));
}
public static ChatResponseMetadata from(ChatResponse chatResponse, Usage usage) {
@@ -217,6 +218,10 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements Cha
return builder.build();
}
private static DefaultUsage getDefaultUsage(CompletionsUsage usage) {
return new DefaultUsage(usage.getPromptTokens(), usage.getCompletionTokens(), usage.getTotalTokens(), usage);
}
public AzureOpenAiChatOptions getDefaultOptions() {
return AzureOpenAiChatOptions.fromOptions(this.defaultOptions);
}
@@ -321,7 +326,7 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements Cha
}
// Accumulate the usage from the previous chat response
CompletionsUsage usage = chatCompletion.getUsage();
Usage currentChatResponseUsage = usage != null ? AzureOpenAiUsage.from(usage) : new EmptyUsage();
Usage currentChatResponseUsage = usage != null ? getDefaultUsage(usage) : new EmptyUsage();
Usage accumulatedUsage = UsageUtils.getCumulativeUsage(currentChatResponseUsage, previousChatResponse);
return toChatResponse(chatCompletion, accumulatedUsage);
}).buffer(2, 1).map(bufferList -> {
@@ -412,7 +417,7 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements Cha
PromptMetadata promptFilterMetadata = generatePromptMetadata(chatCompletions);
Usage currentUsage = null;
if (chatCompletions.getUsage() != null) {
currentUsage = AzureOpenAiUsage.from(chatCompletions);
currentUsage = getDefaultUsage(chatCompletions.getUsage());
}
Usage cumulativeUsage = UsageUtils.getCumulativeUsage(currentUsage, previousChatResponse);
return new ChatResponse(generations, from(chatCompletions, promptFilterMetadata, cumulativeUsage));

View File

@@ -23,11 +23,12 @@ import com.azure.ai.openai.OpenAIClient;
import com.azure.ai.openai.models.EmbeddingItem;
import com.azure.ai.openai.models.Embeddings;
import com.azure.ai.openai.models.EmbeddingsOptions;
import com.azure.ai.openai.models.EmbeddingsUsage;
import io.micrometer.observation.ObservationRegistry;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.azure.openai.metadata.AzureOpenAiEmbeddingUsage;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.document.Document;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.embedding.AbstractEmbeddingModel;
@@ -159,10 +160,14 @@ public class AzureOpenAiEmbeddingModel extends AbstractEmbeddingModel {
private EmbeddingResponse generateEmbeddingResponse(Embeddings embeddings) {
List<Embedding> data = generateEmbeddingList(embeddings.getData());
EmbeddingResponseMetadata metadata = new EmbeddingResponseMetadata();
metadata.setUsage(AzureOpenAiEmbeddingUsage.from(embeddings.getUsage()));
metadata.setUsage(getDefaultUsage(embeddings.getUsage()));
return new EmbeddingResponse(data, metadata);
}
private DefaultUsage getDefaultUsage(EmbeddingsUsage usage) {
return new DefaultUsage(usage.getPromptTokens(), 0, usage.getTotalTokens(), usage);
}
private List<Embedding> generateEmbeddingList(List<EmbeddingItem> nativeData) {
List<Embedding> data = new ArrayList<>();
for (EmbeddingItem nativeDatum : nativeData) {

View File

@@ -1,68 +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.azure.openai.metadata;
import com.azure.ai.openai.models.EmbeddingsUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal Microsoft Azure OpenAI Service} embedding.
*
* @author Thomas Vitale
* @see EmbeddingsUsage
*/
public class AzureOpenAiEmbeddingUsage implements Usage {
private final EmbeddingsUsage usage;
public AzureOpenAiEmbeddingUsage(EmbeddingsUsage usage) {
Assert.notNull(usage, "EmbeddingsUsage must not be null");
this.usage = usage;
}
public static AzureOpenAiEmbeddingUsage from(EmbeddingsUsage usage) {
Assert.notNull(usage, "EmbeddingsUsage must not be null");
return new AzureOpenAiEmbeddingUsage(usage);
}
protected EmbeddingsUsage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
return (long) getUsage().getPromptTokens();
}
@Override
public Long getGenerationTokens() {
return 0L;
}
@Override
public Long getTotalTokens() {
return (long) getUsage().getTotalTokens();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -1,74 +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.azure.openai.metadata;
import com.azure.ai.openai.models.ChatCompletions;
import com.azure.ai.openai.models.CompletionsUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal Microsoft Azure OpenAI Service} chat.
*
* @author John Blum
* @see com.azure.ai.openai.models.CompletionsUsage
* @since 0.7.0
*/
public class AzureOpenAiUsage implements Usage {
private final CompletionsUsage usage;
public AzureOpenAiUsage(CompletionsUsage usage) {
Assert.notNull(usage, "CompletionsUsage must not be null");
this.usage = usage;
}
public static AzureOpenAiUsage from(ChatCompletions chatCompletions) {
Assert.notNull(chatCompletions, "ChatCompletions must not be null");
return from(chatCompletions.getUsage());
}
public static AzureOpenAiUsage from(CompletionsUsage usage) {
return new AzureOpenAiUsage(usage);
}
protected CompletionsUsage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
return (long) getUsage().getPromptTokens();
}
@Override
public Long getGenerationTokens() {
return (long) getUsage().getCompletionTokens();
}
@Override
public Long getTotalTokens() {
return (long) getUsage().getTotalTokens();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -166,7 +166,7 @@ class AzureOpenAiChatModelObservationIT {
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(
ChatModelObservationDocumentation.HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(
ChatModelObservationDocumentation.HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))

View File

@@ -115,7 +115,7 @@ class AzureOpenAiChatModelMetadataTests {
assertThat(usage).isNotNull();
assertThat(usage.getPromptTokens()).isEqualTo(58);
assertThat(usage.getGenerationTokens()).isEqualTo(68);
assertThat(usage.getCompletionTokens()).isEqualTo(68);
assertThat(usage.getTotalTokens()).isEqualTo(126);
}

View File

@@ -547,15 +547,15 @@ public class BedrockProxyChatModel extends AbstractToolCallSupport implements Ch
allGenerations.add(toolCallGeneration);
}
Long promptTokens = response.usage().inputTokens().longValue();
Long generationTokens = response.usage().outputTokens().longValue();
Long totalTokens = response.usage().totalTokens().longValue();
Integer promptTokens = response.usage().inputTokens();
Integer generationTokens = response.usage().outputTokens();
int totalTokens = response.usage().totalTokens();
if (perviousChatResponse != null && perviousChatResponse.getMetadata() != null
&& perviousChatResponse.getMetadata().getUsage() != null) {
promptTokens += perviousChatResponse.getMetadata().getUsage().getPromptTokens();
generationTokens += perviousChatResponse.getMetadata().getUsage().getGenerationTokens();
generationTokens += perviousChatResponse.getMetadata().getUsage().getCompletionTokens();
totalTokens += perviousChatResponse.getMetadata().getUsage().getTotalTokens();
}

View File

@@ -1,63 +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.bedrock.converse.api;
import software.amazon.awssdk.services.bedrockruntime.model.TokenUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for Bedrock Converse API.
*
* @author Christian Tzolov
* @author Wei Jiang
* @since 1.0.0
*/
public class BedrockUsage implements Usage {
public static BedrockUsage from(TokenUsage usage) {
Assert.notNull(usage, "'TokenUsage' must not be null.");
return new BedrockUsage(usage.inputTokens().longValue(), usage.outputTokens().longValue());
}
private final Long inputTokens;
private final Long outputTokens;
protected BedrockUsage(Long inputTokens, Long outputTokens) {
this.inputTokens = inputTokens;
this.outputTokens = outputTokens;
}
@Override
public Long getPromptTokens() {
return this.inputTokens;
}
@Override
public Long getGenerationTokens() {
return this.outputTokens;
}
@Override
public String toString() {
return "BedrockUsage [inputTokens=" + this.inputTokens + ", outputTokens=" + this.outputTokens + "]";
}
}

View File

@@ -122,9 +122,9 @@ public final class ConverseApiUtils {
List<AssistantMessage.ToolCall> toolCalls = new ArrayList<>();
Long promptTokens = 0L;
Long generationTokens = 0L;
Long totalTokens = 0L;
Integer promptTokens = 0;
Integer generationTokens = 0;
Integer totalTokens = 0;
for (ToolUseAggregationEvent.ToolUseEntry toolUseEntry : toolUseAggregationEvent.toolUseEntries()) {
var functionCallId = toolUseEntry.id();
@@ -135,7 +135,7 @@ public final class ConverseApiUtils {
if (toolUseEntry.usage() != null) {
promptTokens += toolUseEntry.usage().getPromptTokens();
generationTokens += toolUseEntry.usage().getGenerationTokens();
generationTokens += toolUseEntry.usage().getCompletionTokens();
totalTokens += toolUseEntry.usage().getTotalTokens();
}
}
@@ -207,9 +207,8 @@ public final class ConverseApiUtils {
Document modelResponseFields = lastAggregation.metadataAggregation().additionalModelResponseFields();
ConverseStreamMetrics metrics = metadataEvent.metrics();
DefaultUsage usage = new DefaultUsage(metadataEvent.usage().inputTokens().longValue(),
metadataEvent.usage().outputTokens().longValue(),
metadataEvent.usage().totalTokens().longValue());
DefaultUsage usage = new DefaultUsage(metadataEvent.usage().inputTokens(),
metadataEvent.usage().outputTokens(), metadataEvent.usage().totalTokens());
var chatResponseMetaData = ChatResponseMetadata.builder().usage(usage).build();
@@ -231,9 +230,9 @@ public final class ConverseApiUtils {
var metadataBuilder = ChatResponseMetadata.builder();
Long promptTokens = perviousChatResponse.getMetadata().getUsage().getPromptTokens();
Long generationTokens = perviousChatResponse.getMetadata().getUsage().getGenerationTokens();
Long totalTokens = perviousChatResponse.getMetadata().getUsage().getTotalTokens();
Integer promptTokens = perviousChatResponse.getMetadata().getUsage().getPromptTokens();
Integer generationTokens = perviousChatResponse.getMetadata().getUsage().getCompletionTokens();
int totalTokens = perviousChatResponse.getMetadata().getUsage().getTotalTokens();
if (chatResponse.getMetadata() != null) {
metadataBuilder.id(chatResponse.getMetadata().getId());
@@ -244,7 +243,7 @@ public final class ConverseApiUtils {
if (chatResponse.getMetadata().getUsage() != null) {
promptTokens = promptTokens + chatResponse.getMetadata().getUsage().getPromptTokens();
generationTokens = generationTokens
+ chatResponse.getMetadata().getUsage().getGenerationTokens();
+ chatResponse.getMetadata().getUsage().getCompletionTokens();
totalTokens = totalTokens + chatResponse.getMetadata().getUsage().getTotalTokens();
}
}
@@ -290,8 +289,8 @@ public final class ConverseApiUtils {
}
else if (event.sdkEventType() == EventType.METADATA) {
ConverseStreamMetadataEvent metadataEvent = (ConverseStreamMetadataEvent) event;
DefaultUsage usage = new DefaultUsage(metadataEvent.usage().inputTokens().longValue(),
metadataEvent.usage().outputTokens().longValue(), metadataEvent.usage().totalTokens().longValue());
DefaultUsage usage = new DefaultUsage(metadataEvent.usage().inputTokens(),
metadataEvent.usage().outputTokens(), metadataEvent.usage().totalTokens());
toolUseEventAggregator.withUsage(usage);
if (!toolUseEventAggregator.isEmpty()) {

View File

@@ -249,11 +249,11 @@ class BedrockConverseChatClientIT {
assertThat(metadata.getUsage().getPromptTokens()).isGreaterThan(500);
assertThat(metadata.getUsage().getPromptTokens()).isLessThan(3500);
assertThat(metadata.getUsage().getGenerationTokens()).isGreaterThan(0);
assertThat(metadata.getUsage().getGenerationTokens()).isLessThan(1500);
assertThat(metadata.getUsage().getCompletionTokens()).isGreaterThan(0);
assertThat(metadata.getUsage().getCompletionTokens()).isLessThan(1500);
assertThat(metadata.getUsage().getTotalTokens())
.isEqualTo(metadata.getUsage().getPromptTokens() + metadata.getUsage().getGenerationTokens());
.isEqualTo(metadata.getUsage().getPromptTokens() + metadata.getUsage().getCompletionTokens());
logger.info("Response: {}", response);
@@ -330,11 +330,11 @@ class BedrockConverseChatClientIT {
assertThat(metadata.getUsage().getPromptTokens()).isGreaterThan(1500);
assertThat(metadata.getUsage().getPromptTokens()).isLessThan(3500);
assertThat(metadata.getUsage().getGenerationTokens()).isGreaterThan(0);
assertThat(metadata.getUsage().getGenerationTokens()).isLessThan(1500);
assertThat(metadata.getUsage().getCompletionTokens()).isGreaterThan(0);
assertThat(metadata.getUsage().getCompletionTokens()).isLessThan(1500);
assertThat(metadata.getUsage().getTotalTokens())
.isEqualTo(metadata.getUsage().getPromptTokens() + metadata.getUsage().getGenerationTokens());
.isEqualTo(metadata.getUsage().getPromptTokens() + metadata.getUsage().getCompletionTokens());
String content = chatResponses.stream()
.filter(cr -> cr.getResult() != null)

View File

@@ -85,7 +85,7 @@ public class BedrockConverseUsageAggregationTests {
assertThat(result.getResult().getOutput().getText()).isSameAs("Response Content Block");
assertThat(result.getMetadata().getUsage().getPromptTokens()).isEqualTo(16);
assertThat(result.getMetadata().getUsage().getGenerationTokens()).isEqualTo(14);
assertThat(result.getMetadata().getUsage().getCompletionTokens()).isEqualTo(14);
assertThat(result.getMetadata().getUsage().getTotalTokens()).isEqualTo(30);
}
@@ -151,7 +151,7 @@ public class BedrockConverseUsageAggregationTests {
.isSameAs(converseResponseFinal.output().message().content().get(0).text());
assertThat(result.getMetadata().getUsage().getPromptTokens()).isEqualTo(445 + 540);
assertThat(result.getMetadata().getUsage().getGenerationTokens()).isEqualTo(119 + 106);
assertThat(result.getMetadata().getUsage().getCompletionTokens()).isEqualTo(119 + 106);
assertThat(result.getMetadata().getUsage().getTotalTokens()).isEqualTo(564 + 646);
}

View File

@@ -77,7 +77,7 @@ class BedrockProxyChatModelIT {
// assertThat(response.getMetadata().getId()).isNotEmpty();
// assertThat(response.getMetadata().getModel()).containsIgnoringCase(model);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}
@@ -93,11 +93,11 @@ class BedrockProxyChatModelIT {
FunctionCallingOptions.builder().model(modelName).build());
ChatResponse response = this.chatModel.call(prompt);
assertThat(response.getResults()).hasSize(1);
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isGreaterThan(0);
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isGreaterThan(0);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isGreaterThan(0);
assertThat(response.getMetadata().getUsage().getTotalTokens())
.isEqualTo(response.getMetadata().getUsage().getPromptTokens()
+ response.getMetadata().getUsage().getGenerationTokens());
+ response.getMetadata().getUsage().getCompletionTokens());
Generation generation = response.getResults().get(0);
assertThat(generation.getOutput().getText()).contains("Blackbeard");
assertThat(generation.getMetadata().getFinishReason()).isEqualTo("end_turn");
@@ -133,11 +133,11 @@ class BedrockProxyChatModelIT {
var referenceTokenUsage = this.chatModel.call(prompt).getMetadata().getUsage();
assertThat(streamingTokenUsage.getPromptTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getGenerationTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getCompletionTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isEqualTo(referenceTokenUsage.getPromptTokens());
assertThat(streamingTokenUsage.getGenerationTokens()).isEqualTo(referenceTokenUsage.getGenerationTokens());
assertThat(streamingTokenUsage.getCompletionTokens()).isEqualTo(referenceTokenUsage.getCompletionTokens());
assertThat(streamingTokenUsage.getTotalTokens()).isEqualTo(referenceTokenUsage.getTotalTokens());
}

View File

@@ -149,7 +149,7 @@ public class BedrockProxyChatModelObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

View File

@@ -1,61 +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.bedrock;
import org.springframework.ai.bedrock.api.AbstractBedrockApi.AmazonBedrockInvocationMetrics;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for Bedrock API.
*
* @author Christian Tzolov
* @since 0.8.0
*/
public class BedrockUsage implements Usage {
private final AmazonBedrockInvocationMetrics usage;
protected BedrockUsage(AmazonBedrockInvocationMetrics usage) {
Assert.notNull(usage, "Bedrock Usage must not be null");
this.usage = usage;
}
public static BedrockUsage from(AmazonBedrockInvocationMetrics usage) {
return new BedrockUsage(usage);
}
protected AmazonBedrockInvocationMetrics getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
return getUsage().inputTokenCount().longValue();
}
@Override
public Long getGenerationTokens() {
return getUsage().outputTokenCount().longValue();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -132,8 +132,7 @@ public class BedrockAnthropic3ChatModel implements ChatModel, StreamingChatModel
}
protected Usage extractUsage(AnthropicChatResponse response) {
return new DefaultUsage(response.usage().inputTokens().longValue(),
response.usage().outputTokens().longValue());
return new DefaultUsage(response.usage().inputTokens(), response.usage().outputTokens());
}
/**

View File

@@ -20,13 +20,14 @@ import java.util.List;
import reactor.core.publisher.Flux;
import org.springframework.ai.bedrock.BedrockUsage;
import org.springframework.ai.bedrock.MessageToPromptConverter;
import org.springframework.ai.bedrock.api.AbstractBedrockApi;
import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi;
import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatRequest;
import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatResponse;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
@@ -80,7 +81,7 @@ public class BedrockCohereChatModel implements ChatModel, StreamingChatModel {
return this.chatApi.chatCompletionStream(this.createRequest(prompt, true)).map(g -> {
if (g.isFinished()) {
String finishReason = g.finishReason().name();
Usage usage = BedrockUsage.from(g.amazonBedrockInvocationMetrics());
Usage usage = getDefaultUsage(g.amazonBedrockInvocationMetrics());
return new ChatResponse(List.of(new Generation(new AssistantMessage(""),
ChatGenerationMetadata.builder().finishReason(finishReason).metadata("usage", usage).build())));
}
@@ -88,6 +89,11 @@ public class BedrockCohereChatModel implements ChatModel, StreamingChatModel {
});
}
private DefaultUsage getDefaultUsage(AbstractBedrockApi.AmazonBedrockInvocationMetrics usage) {
return new DefaultUsage(usage.inputTokenCount().intValue(), usage.outputTokenCount().intValue(),
usage.inputTokenCount().intValue() + usage.outputTokenCount().intValue(), usage);
}
/**
* Test access.
*/

View File

@@ -16,7 +16,9 @@
package org.springframework.ai.bedrock.llama;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import reactor.core.publisher.Flux;
@@ -100,13 +102,22 @@ public class BedrockLlamaChatModel implements ChatModel, StreamingChatModel {
return new Usage() {
@Override
public Long getPromptTokens() {
return response.promptTokenCount().longValue();
public Integer getPromptTokens() {
return response.promptTokenCount();
}
@Override
public Long getGenerationTokens() {
return response.generationTokenCount().longValue();
public Integer getCompletionTokens() {
return response.generationTokenCount();
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", getPromptTokens());
usage.put("completionTokens", getCompletionTokens());
usage.put("totalTokens", getTotalTokens());
return usage;
}
};
}

View File

@@ -16,7 +16,9 @@
package org.springframework.ai.bedrock.titan;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import reactor.core.publisher.Flux;
@@ -140,13 +142,22 @@ public class BedrockTitanChatModel implements ChatModel, StreamingChatModel {
return new Usage() {
@Override
public Long getPromptTokens() {
return response.inputTextTokenCount().longValue();
public Integer getPromptTokens() {
return response.inputTextTokenCount();
}
@Override
public Long getGenerationTokens() {
return response.totalOutputTextTokenCount().longValue();
public Integer getCompletionTokens() {
return response.totalOutputTextTokenCount();
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", getPromptTokens());
usage.put("completionTokens", getCompletionTokens());
usage.put("totalTokens", getTotalTokens());
return usage;
}
};
}

View File

@@ -36,6 +36,7 @@ import org.springframework.ai.chat.messages.MessageType;
import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.model.AbstractToolCallSupport;
import org.springframework.ai.chat.model.ChatModel;
@@ -60,7 +61,6 @@ import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletionMessage.Role;
import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletionMessage.ToolCall;
import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletionRequest;
import org.springframework.ai.minimax.api.MiniMaxApiConstants;
import org.springframework.ai.minimax.metadata.MiniMaxUsage;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackResolver;
@@ -388,13 +388,17 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
Assert.notNull(result, "MiniMax ChatCompletionResult must not be null");
return ChatResponseMetadata.builder()
.id(result.id() != null ? result.id() : "")
.usage(result.usage() != null ? MiniMaxUsage.from(result.usage()) : new EmptyUsage())
.usage(result.usage() != null ? getDefaultUsage(result.usage()) : new EmptyUsage())
.model(result.model() != null ? result.model() : "")
.keyValue("created", result.created() != null ? result.created() : 0L)
.keyValue("system-fingerprint", result.systemFingerprint() != null ? result.systemFingerprint() : "")
.build();
}
private DefaultUsage getDefaultUsage(MiniMaxApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
private Generation buildGeneration(ChatCompletionMessage message, ChatCompletionFinishReason completionFinishReason,
Map<String, Object> metadata) {
if (message == null || message.role() == Role.TOOL) {

View File

@@ -23,6 +23,7 @@ import io.micrometer.observation.ObservationRegistry;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.document.Document;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.embedding.AbstractEmbeddingModel;
@@ -37,7 +38,6 @@ import org.springframework.ai.embedding.observation.EmbeddingModelObservationCon
import org.springframework.ai.embedding.observation.EmbeddingModelObservationDocumentation;
import org.springframework.ai.minimax.api.MiniMaxApi;
import org.springframework.ai.minimax.api.MiniMaxApiConstants;
import org.springframework.ai.minimax.metadata.MiniMaxUsage;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.lang.Nullable;
@@ -171,8 +171,7 @@ public class MiniMaxEmbeddingModel extends AbstractEmbeddingModel {
return new EmbeddingResponse(List.of());
}
var metadata = new EmbeddingResponseMetadata(apiRequest.model(),
MiniMaxUsage.from(new MiniMaxApi.Usage(0, 0, apiEmbeddingResponse.totalTokens())));
var metadata = new EmbeddingResponseMetadata(apiRequest.model(), getDefaultUsage(apiEmbeddingResponse));
List<Embedding> embeddings = new ArrayList<>();
for (int i = 0; i < apiEmbeddingResponse.vectors().size(); i++) {
@@ -185,6 +184,10 @@ public class MiniMaxEmbeddingModel extends AbstractEmbeddingModel {
});
}
private DefaultUsage getDefaultUsage(MiniMaxApi.EmbeddingList apiEmbeddingList) {
return new DefaultUsage(0, 0, apiEmbeddingList.totalTokens());
}
/**
* Merge runtime and default {@link EmbeddingOptions} to compute the final options to
* use in the request.

View File

@@ -1,71 +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.minimax.metadata;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.minimax.api.MiniMaxApi;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal MiniMax}.
*
* @author Thomas Vitale
*/
public class MiniMaxUsage implements Usage {
private final MiniMaxApi.Usage usage;
protected MiniMaxUsage(MiniMaxApi.Usage usage) {
Assert.notNull(usage, "MiniMax Usage must not be null");
this.usage = usage;
}
public static MiniMaxUsage from(MiniMaxApi.Usage usage) {
return new MiniMaxUsage(usage);
}
protected MiniMaxApi.Usage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
Integer promptTokens = getUsage().promptTokens();
return promptTokens != null ? promptTokens.longValue() : 0;
}
@Override
public Long getGenerationTokens() {
Integer generationTokens = getUsage().completionTokens();
return generationTokens != null ? generationTokens.longValue() : 0;
}
@Override
public Long getTotalTokens() {
Integer totalTokens = getUsage().totalTokens();
if (totalTokens != null) {
return totalTokens.longValue();
}
return getPromptTokens() + getGenerationTokens();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -148,7 +148,7 @@ public class MiniMaxChatModelObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

View File

@@ -38,6 +38,7 @@ import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.metadata.UsageUtils;
import org.springframework.ai.chat.model.AbstractToolCallSupport;
@@ -59,7 +60,6 @@ import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionMessage;
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionMessage.ChatCompletionFunction;
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionMessage.ToolCall;
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionRequest;
import org.springframework.ai.mistralai.metadata.MistralAiUsage;
import org.springframework.ai.model.Media;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.model.function.FunctionCallback;
@@ -154,7 +154,7 @@ public class MistralAiChatModel extends AbstractToolCallSupport implements ChatM
public static ChatResponseMetadata from(MistralAiApi.ChatCompletion result) {
Assert.notNull(result, "Mistral AI ChatCompletion must not be null");
MistralAiUsage usage = MistralAiUsage.from(result.usage());
DefaultUsage usage = getDefaultUsage(result.usage());
return ChatResponseMetadata.builder()
.id(result.id())
.model(result.model())
@@ -173,6 +173,10 @@ public class MistralAiChatModel extends AbstractToolCallSupport implements ChatM
.build();
}
private static DefaultUsage getDefaultUsage(MistralAiApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
@Override
public ChatResponse call(Prompt prompt) {
return this.internalCall(prompt, null);
@@ -214,7 +218,7 @@ public class MistralAiChatModel extends AbstractToolCallSupport implements ChatM
return buildGeneration(choice, metadata);
}).toList();
MistralAiUsage usage = MistralAiUsage.from(completionEntity.getBody().usage());
DefaultUsage usage = getDefaultUsage(completionEntity.getBody().usage());
Usage cumulativeUsage = UsageUtils.getCumulativeUsage(usage, previousChatResponse);
ChatResponse chatResponse = new ChatResponse(generations,
from(completionEntity.getBody(), cumulativeUsage));
@@ -287,7 +291,7 @@ public class MistralAiChatModel extends AbstractToolCallSupport implements ChatM
// @formatter:on
if (chatCompletion2.usage() != null) {
MistralAiUsage usage = MistralAiUsage.from(chatCompletion2.usage());
DefaultUsage usage = getDefaultUsage(chatCompletion2.usage());
Usage cumulativeUsage = UsageUtils.getCumulativeUsage(usage, previousChatResponse);
return new ChatResponse(generations, from(chatCompletion2, cumulativeUsage));
}

View File

@@ -22,6 +22,7 @@ import io.micrometer.observation.ObservationRegistry;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.document.Document;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.embedding.AbstractEmbeddingModel;
@@ -36,7 +37,6 @@ import org.springframework.ai.embedding.observation.EmbeddingModelObservationCon
import org.springframework.ai.embedding.observation.EmbeddingModelObservationConvention;
import org.springframework.ai.embedding.observation.EmbeddingModelObservationDocumentation;
import org.springframework.ai.mistralai.api.MistralAiApi;
import org.springframework.ai.mistralai.metadata.MistralAiUsage;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.retry.support.RetryTemplate;
@@ -131,7 +131,7 @@ public class MistralAiEmbeddingModel extends AbstractEmbeddingModel {
}
var metadata = new EmbeddingResponseMetadata(apiEmbeddingResponse.model(),
MistralAiUsage.from(apiEmbeddingResponse.usage()));
getDefaultUsage(apiEmbeddingResponse.usage()));
var embeddings = apiEmbeddingResponse.data()
.stream()
@@ -146,6 +146,10 @@ public class MistralAiEmbeddingModel extends AbstractEmbeddingModel {
});
}
private DefaultUsage getDefaultUsage(MistralAiApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
@SuppressWarnings("unchecked")
private MistralAiApi.EmbeddingRequest<List<String>> createRequest(EmbeddingRequest request) {
var embeddingRequest = new MistralAiApi.EmbeddingRequest<>(request.getInstructions(),

View File

@@ -1,67 +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.mistralai.metadata;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.mistralai.api.MistralAiApi;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal Mistral AI}.
*
* @author Thomas Vitale
* @since 1.0.0
* @see <a href="https://docs.mistral.ai/api/">Chat Completion API</a>
*/
public class MistralAiUsage implements Usage {
private final MistralAiApi.Usage usage;
protected MistralAiUsage(MistralAiApi.Usage usage) {
Assert.notNull(usage, "Mistral AI Usage must not be null");
this.usage = usage;
}
public static MistralAiUsage from(MistralAiApi.Usage usage) {
return new MistralAiUsage(usage);
}
protected MistralAiApi.Usage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
return getUsage().promptTokens().longValue();
}
@Override
public Long getGenerationTokens() {
return getUsage().completionTokens().longValue();
}
@Override
public Long getTotalTokens() {
return getUsage().totalTokens().longValue();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -305,7 +305,7 @@ class MistralAiChatClientIT {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(model);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}

View File

@@ -158,7 +158,7 @@ public class MistralAiChatModelObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

View File

@@ -35,6 +35,7 @@ import org.springframework.ai.chat.messages.MessageType;
import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.metadata.UsageUtils;
@@ -65,7 +66,6 @@ import org.springframework.ai.moonshot.api.MoonshotApi.ChatCompletionMessage.Too
import org.springframework.ai.moonshot.api.MoonshotApi.ChatCompletionRequest;
import org.springframework.ai.moonshot.api.MoonshotApi.FunctionTool;
import org.springframework.ai.moonshot.api.MoonshotConstants;
import org.springframework.ai.moonshot.metadata.MoonshotUsage;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
@@ -226,7 +226,7 @@ public class MoonshotChatModel extends AbstractToolCallSupport implements ChatMo
return buildGeneration(choice, metadata);
}).toList();
MoonshotApi.Usage usage = completionEntity.getBody().usage();
Usage currentUsage = (usage != null) ? MoonshotUsage.from(usage) : new EmptyUsage();
Usage currentUsage = (usage != null) ? getDefaultUsage(usage) : new EmptyUsage();
Usage cumulativeUsage = UsageUtils.getCumulativeUsage(currentUsage, previousChatResponse);
ChatResponse chatResponse = new ChatResponse(generations,
from(completionEntity.getBody(), cumulativeUsage));
@@ -247,6 +247,10 @@ public class MoonshotChatModel extends AbstractToolCallSupport implements ChatMo
return response;
}
private DefaultUsage getDefaultUsage(MoonshotApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
@Override
public ChatOptions getDefaultOptions() {
return this.defaultOptions.copy();
@@ -302,7 +306,7 @@ public class MoonshotChatModel extends AbstractToolCallSupport implements ChatMo
return buildGeneration(choice, metadata);
}).toList();
MoonshotApi.Usage usage = chatCompletion2.usage();
Usage currentUsage = (usage != null) ? MoonshotUsage.from(usage) : new EmptyUsage();
Usage currentUsage = (usage != null) ? getDefaultUsage(usage) : new EmptyUsage();
Usage cumulativeUsage = UsageUtils.getCumulativeUsage(currentUsage, previousChatResponse);
return new ChatResponse(generations, from(chatCompletion2, cumulativeUsage));
@@ -336,7 +340,7 @@ public class MoonshotChatModel extends AbstractToolCallSupport implements ChatMo
Assert.notNull(result, "Moonshot ChatCompletionResult must not be null");
return ChatResponseMetadata.builder()
.id(result.id() != null ? result.id() : "")
.usage(result.usage() != null ? MoonshotUsage.from(result.usage()) : new EmptyUsage())
.usage(result.usage() != null ? getDefaultUsage(result.usage()) : new EmptyUsage())
.model(result.model() != null ? result.model() : "")
.keyValue("created", result.created() != null ? result.created() : 0L)
.build();

View File

@@ -1,65 +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.moonshot.metadata;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.moonshot.api.MoonshotApi;
import org.springframework.util.Assert;
/**
* Represents the usage of a Moonshot model.
*
* @author Geng Rong
*/
public class MoonshotUsage implements Usage {
private final MoonshotApi.Usage usage;
protected MoonshotUsage(MoonshotApi.Usage usage) {
Assert.notNull(usage, "Moonshot Usage must not be null");
this.usage = usage;
}
public static MoonshotUsage from(MoonshotApi.Usage usage) {
return new MoonshotUsage(usage);
}
protected MoonshotApi.Usage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
return getUsage().promptTokens().longValue();
}
@Override
public Long getGenerationTokens() {
return getUsage().completionTokens().longValue();
}
@Override
public Long getTotalTokens() {
return getUsage().totalTokens().longValue();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -150,7 +150,7 @@ public class MoonshotChatModelObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

View File

@@ -21,6 +21,7 @@ import java.util.Base64;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.Set;
import io.micrometer.observation.Observation;
@@ -60,7 +61,6 @@ import org.springframework.ai.ollama.api.OllamaOptions;
import org.springframework.ai.ollama.management.ModelManagementOptions;
import org.springframework.ai.ollama.management.OllamaModelManager;
import org.springframework.ai.ollama.management.PullModelStrategy;
import org.springframework.ai.ollama.metadata.OllamaChatUsage;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.StringUtils;
@@ -132,10 +132,10 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode
static ChatResponseMetadata from(OllamaApi.ChatResponse response, ChatResponse previousChatResponse) {
Assert.notNull(response, "OllamaApi.ChatResponse must not be null");
OllamaChatUsage newUsage = OllamaChatUsage.from(response);
Long promptTokens = newUsage.getPromptTokens();
Long generationTokens = newUsage.getGenerationTokens();
Long totalTokens = newUsage.getTotalTokens();
DefaultUsage newUsage = getDefaultUsage(response);
Integer promptTokens = newUsage.getPromptTokens();
Integer generationTokens = newUsage.getCompletionTokens();
int totalTokens = newUsage.getTotalTokens();
Duration evalDuration = response.getEvalDuration();
Duration promptEvalDuration = response.getPromptEvalDuration();
@@ -158,7 +158,7 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode
}
if (previousChatResponse.getMetadata().getUsage() != null) {
promptTokens += previousChatResponse.getMetadata().getUsage().getPromptTokens();
generationTokens += previousChatResponse.getMetadata().getUsage().getGenerationTokens();
generationTokens += previousChatResponse.getMetadata().getUsage().getCompletionTokens();
totalTokens += previousChatResponse.getMetadata().getUsage().getTotalTokens();
}
}
@@ -170,7 +170,7 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode
.model(response.model())
.keyValue(METADATA_CREATED_AT, response.createdAt())
.keyValue(METADATA_EVAL_DURATION, evalDuration)
.keyValue(METADATA_EVAL_COUNT, aggregatedUsage.getGenerationTokens().intValue())
.keyValue(METADATA_EVAL_COUNT, aggregatedUsage.getCompletionTokens().intValue())
.keyValue(METADATA_LOAD_DURATION, loadDuration)
.keyValue(METADATA_PROMPT_EVAL_DURATION, promptEvalDuration)
.keyValue(METADATA_PROMPT_EVAL_COUNT, aggregatedUsage.getPromptTokens().intValue())
@@ -179,6 +179,11 @@ public class OllamaChatModel extends AbstractToolCallSupport implements ChatMode
.build();
}
private static DefaultUsage getDefaultUsage(OllamaApi.ChatResponse response) {
return new DefaultUsage(Optional.ofNullable(response.promptEvalCount()).orElse(0),
Optional.ofNullable(response.evalCount()).orElse(0));
}
@Override
public ChatResponse call(Prompt prompt) {
return this.internalCall(prompt, null);

View File

@@ -18,12 +18,14 @@ package org.springframework.ai.ollama;
import java.time.Duration;
import java.util.List;
import java.util.Optional;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.regex.Matcher;
import java.util.regex.Pattern;
import io.micrometer.observation.ObservationRegistry;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.AbstractEmbeddingModel;
import org.springframework.ai.embedding.Embedding;
@@ -45,7 +47,6 @@ import org.springframework.ai.ollama.api.OllamaOptions;
import org.springframework.ai.ollama.management.ModelManagementOptions;
import org.springframework.ai.ollama.management.OllamaModelManager;
import org.springframework.ai.ollama.management.PullModelStrategy;
import org.springframework.ai.ollama.metadata.OllamaEmbeddingUsage;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
@@ -126,7 +127,7 @@ public class OllamaEmbeddingModel extends AbstractEmbeddingModel {
.toList();
EmbeddingResponseMetadata embeddingResponseMetadata = new EmbeddingResponseMetadata(response.model(),
OllamaEmbeddingUsage.from(response));
getDefaultUsage(response));
EmbeddingResponse embeddingResponse = new EmbeddingResponse(embeddings, embeddingResponseMetadata);
@@ -136,6 +137,10 @@ public class OllamaEmbeddingModel extends AbstractEmbeddingModel {
});
}
private DefaultUsage getDefaultUsage(OllamaApi.EmbeddingsResponse response) {
return new DefaultUsage(Optional.ofNullable(response.promptEvalCount()).orElse(0), 0);
}
/**
* Package access for testing.
*/

View File

@@ -1,61 +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.ollama.metadata;
import java.util.Optional;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.ollama.api.OllamaApi;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal Ollama}
*
* @see Usage
* @author Fu Cheng
*/
public class OllamaChatUsage implements Usage {
protected static final String AI_USAGE_STRING = "{ promptTokens: %1$d, generationTokens: %2$d, totalTokens: %3$d }";
private final OllamaApi.ChatResponse response;
public OllamaChatUsage(OllamaApi.ChatResponse response) {
this.response = response;
}
public static OllamaChatUsage from(OllamaApi.ChatResponse response) {
Assert.notNull(response, "OllamaApi.ChatResponse must not be null");
return new OllamaChatUsage(response);
}
@Override
public Long getPromptTokens() {
return Optional.ofNullable(this.response.promptEvalCount()).map(Integer::longValue).orElse(0L);
}
@Override
public Long getGenerationTokens() {
return Optional.ofNullable(this.response.evalCount()).map(Integer::longValue).orElse(0L);
}
@Override
public String toString() {
return AI_USAGE_STRING.formatted(getPromptTokens(), getGenerationTokens(), getTotalTokens());
}
}

View File

@@ -1,61 +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.ollama.metadata;
import java.util.Optional;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.ollama.api.OllamaApi.EmbeddingsResponse;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal Ollama} embeddings.
*
* @see Usage
* @author Christian Tzolov
*/
public class OllamaEmbeddingUsage implements Usage {
protected static final String AI_USAGE_STRING = "{ promptTokens: %1$d, generationTokens: %2$d, totalTokens: %3$d }";
private Long promptTokens;
public OllamaEmbeddingUsage(EmbeddingsResponse response) {
this.promptTokens = Optional.ofNullable(response.promptEvalCount()).map(Integer::longValue).orElse(0L);
}
public static OllamaEmbeddingUsage from(EmbeddingsResponse response) {
Assert.notNull(response, "OllamaApi.EmbeddingsResponse must not be null");
return new OllamaEmbeddingUsage(response);
}
@Override
public Long getPromptTokens() {
return this.promptTokens;
}
@Override
public Long getGenerationTokens() {
return 0L;
}
@Override
public String toString() {
return AI_USAGE_STRING.formatted(getPromptTokens(), getGenerationTokens(), getTotalTokens());
}
}

View File

@@ -139,7 +139,7 @@ class OllamaChatModelIT extends BaseOllamaIT {
assertThat(usage).isNotNull();
assertThat(usage.getPromptTokens()).isPositive();
assertThat(usage.getGenerationTokens()).isPositive();
assertThat(usage.getCompletionTokens()).isPositive();
assertThat(usage.getTotalTokens()).isPositive();
}

View File

@@ -147,7 +147,7 @@ public class OllamaChatModelObservationIT extends BaseOllamaIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

View File

@@ -100,7 +100,7 @@ public class OllamaChatModelTests {
ChatResponse previousChatResponse = ChatResponse.builder()
.generations(List.of())
.metadata(ChatResponseMetadata.builder()
.usage(new DefaultUsage(66L, 99L))
.usage(new DefaultUsage(66, 99))
.keyValue("eval-duration", Duration.ofSeconds(2))
.keyValue("prompt-eval-duration", Duration.ofSeconds(2))
.build())
@@ -108,7 +108,7 @@ public class OllamaChatModelTests {
ChatResponseMetadata metadata = OllamaChatModel.from(response, previousChatResponse);
assertThat(metadata.getUsage()).isEqualTo(new DefaultUsage(808L + 66L, 101L + 99L));
assertThat(metadata.getUsage()).isEqualTo(new DefaultUsage(808 + 66, 101 + 99));
assertEquals(Duration.ofNanos(evalDuration).plus(Duration.ofSeconds(2)), metadata.get("eval-duration"));
assertEquals((evalCount + 99), (Integer) metadata.get("eval-count"));

View File

@@ -40,6 +40,7 @@ import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.metadata.Usage;
@@ -71,7 +72,6 @@ import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.MediaCo
import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.ToolCall;
import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest;
import org.springframework.ai.openai.api.common.OpenAiApiConstants;
import org.springframework.ai.openai.metadata.OpenAiUsage;
import org.springframework.ai.openai.metadata.support.OpenAiResponseHeaderExtractor;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.core.io.ByteArrayResource;
@@ -267,7 +267,7 @@ public class OpenAiChatModel extends AbstractToolCallSupport implements ChatMode
RateLimit rateLimit = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity);
// Current usage
OpenAiApi.Usage usage = completionEntity.getBody().usage();
Usage currentChatResponseUsage = usage != null ? OpenAiUsage.from(usage) : new EmptyUsage();
Usage currentChatResponseUsage = usage != null ? getDefaultUsage(usage) : new EmptyUsage();
Usage accumulatedUsage = UsageUtils.getCumulativeUsage(currentChatResponseUsage, previousChatResponse);
ChatResponse chatResponse = new ChatResponse(generations,
from(completionEntity.getBody(), rateLimit, accumulatedUsage));
@@ -352,7 +352,7 @@ public class OpenAiChatModel extends AbstractToolCallSupport implements ChatMode
}).toList();
// @formatter:on
OpenAiApi.Usage usage = chatCompletion2.usage();
Usage currentChatResponseUsage = usage != null ? OpenAiUsage.from(usage) : new EmptyUsage();
Usage currentChatResponseUsage = usage != null ? getDefaultUsage(usage) : new EmptyUsage();
Usage accumulatedUsage = UsageUtils.getCumulativeUsage(currentChatResponseUsage,
previousChatResponse);
return new ChatResponse(generations, from(chatCompletion2, null, accumulatedUsage));
@@ -501,6 +501,10 @@ public class OpenAiChatModel extends AbstractToolCallSupport implements ChatMode
chunk.systemFingerprint(), "chat.completion", chunk.usage());
}
private DefaultUsage getDefaultUsage(OpenAiApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
/**
* Accessible for testing.
*/

View File

@@ -22,6 +22,7 @@ import io.micrometer.observation.ObservationRegistry;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.document.Document;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.embedding.AbstractEmbeddingModel;
@@ -38,7 +39,6 @@ import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.api.OpenAiApi.EmbeddingList;
import org.springframework.ai.openai.api.common.OpenAiApiConstants;
import org.springframework.ai.openai.metadata.OpenAiUsage;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.lang.Nullable;
import org.springframework.retry.support.RetryTemplate;
@@ -168,7 +168,7 @@ public class OpenAiEmbeddingModel extends AbstractEmbeddingModel {
}
var metadata = new EmbeddingResponseMetadata(apiEmbeddingResponse.model(),
OpenAiUsage.from(apiEmbeddingResponse.usage()));
getDefaultUsage(apiEmbeddingResponse.usage()));
List<Embedding> embeddings = apiEmbeddingResponse.data()
.stream()
@@ -183,6 +183,10 @@ public class OpenAiEmbeddingModel extends AbstractEmbeddingModel {
});
}
private DefaultUsage getDefaultUsage(OpenAiApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
private OpenAiApi.EmbeddingRequest<List<String>> createRequest(EmbeddingRequest request,
OpenAiEmbeddingOptions requestOptions) {
return new OpenAiApi.EmbeddingRequest<>(request.getInstructions(), requestOptions.getModel(),

View File

@@ -1,166 +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.Usage;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal OpenAI}.
*
* @author John Blum
* @author Thomas Vitale
* @author David Frizelle
* @author Christian Tzolov
* @since 0.7.0
* @see <a href=
* "https://platform.openai.com/docs/api-reference/completions/object">Completion
* Object</a>
*/
public class OpenAiUsage implements Usage {
private final OpenAiApi.Usage usage;
protected OpenAiUsage(OpenAiApi.Usage usage) {
Assert.notNull(usage, "OpenAI Usage must not be null");
this.usage = usage;
}
public static OpenAiUsage from(OpenAiApi.Usage usage) {
return new OpenAiUsage(usage);
}
protected OpenAiApi.Usage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
Integer promptTokens = getUsage().promptTokens();
return promptTokens != null ? promptTokens.longValue() : 0;
}
@Override
public Long getGenerationTokens() {
Integer generationTokens = getUsage().completionTokens();
return generationTokens != null ? generationTokens.longValue() : 0;
}
@Override
public Long getTotalTokens() {
Integer totalTokens = getUsage().totalTokens();
if (totalTokens != null) {
return totalTokens.longValue();
}
else {
return getPromptTokens() + getGenerationTokens();
}
}
/**
* @deprecated Use {@link #getPromptTokensDetails()} instead.
*/
@Deprecated
public Long getPromptTokensDetailsCachedTokens() {
OpenAiApi.Usage.PromptTokensDetails promptTokenDetails = getUsage().promptTokensDetails();
Integer cachedTokens = promptTokenDetails != null ? promptTokenDetails.cachedTokens() : null;
return cachedTokens != null ? cachedTokens.longValue() : 0;
}
public PromptTokensDetails getPromptTokensDetails() {
var details = getUsage().promptTokensDetails();
if (details == null) {
return new PromptTokensDetails(0, 0);
}
return new PromptTokensDetails(valueOrZero(details.audioTokens()), valueOrZero(details.cachedTokens()));
}
/**
* @deprecated Use {@link #getCompletionTokenDetails()} instead.
*/
@Deprecated
public Long getReasoningTokens() {
OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails();
Integer reasoningTokens = completionTokenDetails != null ? completionTokenDetails.reasoningTokens() : null;
return reasoningTokens != null ? reasoningTokens.longValue() : 0;
}
/**
* @deprecated Use {@link #getCompletionTokenDetails()} instead.
*/
@Deprecated
public Long getAcceptedPredictionTokens() {
OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails();
Integer acceptedPredictionTokens = completionTokenDetails != null
? completionTokenDetails.acceptedPredictionTokens() : null;
return acceptedPredictionTokens != null ? acceptedPredictionTokens.longValue() : 0;
}
/**
* @deprecated Use {@link #getCompletionTokenDetails()} instead.
*/
@Deprecated
public Long getAudioTokens() {
OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails();
Integer audioTokens = completionTokenDetails != null ? completionTokenDetails.audioTokens() : null;
return audioTokens != null ? audioTokens.longValue() : 0;
}
/**
* @deprecated Use {@link #getCompletionTokenDetails()} instead.
*/
@Deprecated
public Long getRejectedPredictionTokens() {
OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails();
Integer rejectedPredictionTokens = completionTokenDetails != null
? completionTokenDetails.rejectedPredictionTokens() : null;
return rejectedPredictionTokens != null ? rejectedPredictionTokens.longValue() : 0;
}
public CompletionTokenDetails getCompletionTokenDetails() {
var details = getUsage().completionTokenDetails();
if (details == null) {
return new CompletionTokenDetails(0, 0, 0, 0);
}
return new CompletionTokenDetails(valueOrZero(details.reasoningTokens()),
valueOrZero(details.acceptedPredictionTokens()), valueOrZero(details.audioTokens()),
valueOrZero(details.rejectedPredictionTokens()));
}
@Override
public String toString() {
return getUsage().toString();
}
private int valueOrZero(Integer value) {
return value != null ? value : 0;
}
public record PromptTokensDetails(// @formatter:off
Integer audioTokens,
Integer cachedTokens) {
}
public record CompletionTokenDetails(
Integer reasoningTokens,
Integer acceptedPredictionTokens,
Integer audioTokens,
Integer rejectedPredictionTokens) { // @formatter:on
}
}

View File

@@ -214,12 +214,12 @@ public class OpenAiChatModelIT extends AbstractIT {
var referenceTokenUsage = this.chatModel.call(prompt).getMetadata().getUsage();
assertThat(streamingTokenUsage.getPromptTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getGenerationTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getCompletionTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isCloseTo(referenceTokenUsage.getPromptTokens(),
Percentage.withPercentage(25));
assertThat(streamingTokenUsage.getGenerationTokens()).isCloseTo(referenceTokenUsage.getGenerationTokens(),
assertThat(streamingTokenUsage.getCompletionTokens()).isCloseTo(referenceTokenUsage.getCompletionTokens(),
Percentage.withPercentage(25));
assertThat(streamingTokenUsage.getTotalTokens()).isCloseTo(referenceTokenUsage.getTotalTokens(),
Percentage.withPercentage(25));
@@ -413,9 +413,9 @@ public class OpenAiChatModelIT extends AbstractIT {
assertThat(usage).isNotNull();
assertThat(usage).isNotInstanceOf(EmptyUsage.class);
assertThat(usage).isInstanceOf(DefaultUsage.class);
assertThat(usage.getPromptTokens()).isGreaterThan(450L).isLessThan(600L);
assertThat(usage.getGenerationTokens()).isGreaterThan(230L).isLessThan(360L);
assertThat(usage.getTotalTokens()).isGreaterThan(680L).isLessThan(900L);
assertThat(usage.getPromptTokens()).isGreaterThan(450).isLessThan(600);
assertThat(usage.getCompletionTokens()).isGreaterThan(230).isLessThan(360);
assertThat(usage.getTotalTokens()).isGreaterThan(680).isLessThan(900);
}
@Test
@@ -442,9 +442,9 @@ public class OpenAiChatModelIT extends AbstractIT {
assertThat(usage).isNotNull();
assertThat(usage).isNotInstanceOf(EmptyUsage.class);
assertThat(usage).isInstanceOf(DefaultUsage.class);
assertThat(usage.getPromptTokens()).isGreaterThan(450L).isLessThan(600L);
assertThat(usage.getGenerationTokens()).isGreaterThan(230L).isLessThan(360L);
assertThat(usage.getTotalTokens()).isGreaterThan(680L).isLessThan(960L);
assertThat(usage.getPromptTokens()).isGreaterThan(450).isLessThan(600);
assertThat(usage.getCompletionTokens()).isGreaterThan(230).isLessThan(360);
assertThat(usage.getTotalTokens()).isGreaterThan(680).isLessThan(960);
}
@ParameterizedTest(name = "{0} : {displayName} ")
@@ -596,7 +596,7 @@ public class OpenAiChatModelIT extends AbstractIT {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(model);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}

View File

@@ -152,7 +152,7 @@ public class OpenAiChatModelObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

View File

@@ -88,7 +88,7 @@ public class OpenAiChatModelWithChatResponseMetadataTests {
assertThat(usage).isNotNull();
assertThat(usage.getPromptTokens()).isEqualTo(9L);
assertThat(usage.getGenerationTokens()).isEqualTo(12L);
assertThat(usage.getCompletionTokens()).isEqualTo(12L);
assertThat(usage.getTotalTokens()).isEqualTo(21L);
RateLimit rateLimit = chatResponseMetadata.getRateLimit();

View File

@@ -126,11 +126,11 @@ class GroqWithOpenAiChatModelIT {
var referenceTokenUsage = this.chatModel.call(prompt).getMetadata().getUsage();
assertThat(streamingTokenUsage.getPromptTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getGenerationTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getCompletionTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isEqualTo(referenceTokenUsage.getPromptTokens());
assertThat(streamingTokenUsage.getGenerationTokens()).isEqualTo(referenceTokenUsage.getGenerationTokens());
assertThat(streamingTokenUsage.getCompletionTokens()).isEqualTo(referenceTokenUsage.getCompletionTokens());
assertThat(streamingTokenUsage.getTotalTokens()).isEqualTo(referenceTokenUsage.getTotalTokens());
}
@@ -371,7 +371,7 @@ class GroqWithOpenAiChatModelIT {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(model);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}

View File

@@ -125,11 +125,11 @@ class MistralWithOpenAiChatModelIT {
var referenceTokenUsage = this.chatModel.call(prompt).getMetadata().getUsage();
assertThat(streamingTokenUsage.getPromptTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getGenerationTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getCompletionTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isEqualTo(referenceTokenUsage.getPromptTokens());
assertThat(streamingTokenUsage.getGenerationTokens()).isEqualTo(referenceTokenUsage.getGenerationTokens());
assertThat(streamingTokenUsage.getCompletionTokens()).isEqualTo(referenceTokenUsage.getCompletionTokens());
assertThat(streamingTokenUsage.getTotalTokens()).isEqualTo(referenceTokenUsage.getTotalTokens());
}
@@ -376,7 +376,7 @@ class MistralWithOpenAiChatModelIT {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(model);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}

View File

@@ -122,11 +122,11 @@ class NvidiaWithOpenAiChatModelIT {
var referenceTokenUsage = this.chatModel.call(prompt).getMetadata().getUsage();
assertThat(streamingTokenUsage.getPromptTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getGenerationTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getCompletionTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isEqualTo(referenceTokenUsage.getPromptTokens());
assertThat(streamingTokenUsage.getGenerationTokens()).isEqualTo(referenceTokenUsage.getGenerationTokens());
assertThat(streamingTokenUsage.getCompletionTokens()).isEqualTo(referenceTokenUsage.getCompletionTokens());
assertThat(streamingTokenUsage.getTotalTokens()).isEqualTo(referenceTokenUsage.getTotalTokens());
}
@@ -305,7 +305,7 @@ class NvidiaWithOpenAiChatModelIT {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(DEFAULT_NVIDIA_MODEL);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}

View File

@@ -143,11 +143,11 @@ class OllamaWithOpenAiChatModelIT {
var referenceTokenUsage = this.chatModel.call(prompt).getMetadata().getUsage();
// assertThat(streamingTokenUsage.getPromptTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getGenerationTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getCompletionTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isEqualTo(referenceTokenUsage.getPromptTokens());
assertThat(streamingTokenUsage.getGenerationTokens()).isEqualTo(referenceTokenUsage.getGenerationTokens());
assertThat(streamingTokenUsage.getCompletionTokens()).isEqualTo(referenceTokenUsage.getCompletionTokens());
assertThat(streamingTokenUsage.getTotalTokens()).isEqualTo(referenceTokenUsage.getTotalTokens());
}
@@ -400,7 +400,7 @@ class OllamaWithOpenAiChatModelIT {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(model);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}

View File

@@ -139,12 +139,12 @@ class PerplexityWithOpenAiChatModelIT {
var referenceTokenUsage = this.chatModel.call(prompt).getMetadata().getUsage();
assertThat(streamingTokenUsage.getPromptTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getGenerationTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getCompletionTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isEqualTo(referenceTokenUsage.getPromptTokens());
assertThat(streamingTokenUsage.getGenerationTokens())
.isGreaterThanOrEqualTo(referenceTokenUsage.getGenerationTokens());
assertThat(streamingTokenUsage.getCompletionTokens())
.isGreaterThanOrEqualTo(referenceTokenUsage.getCompletionTokens());
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThanOrEqualTo(referenceTokenUsage.getTotalTokens());
}
@@ -315,7 +315,7 @@ class PerplexityWithOpenAiChatModelIT {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(DEFAULT_PERPLEXITY_MODEL);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
}

View File

@@ -18,129 +18,142 @@ package org.springframework.ai.openai.metadata;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.openai.api.OpenAiApi;
import static org.assertj.core.api.Assertions.assertThat;
/**
* Unit tests for {@link OpenAiUsage}.
* Unit tests for OpenAI usage data.
*
* @author Thomas Vitale
* @author Christian Tzolov
* @author Ilayaperumal Gopinathan
*/
class OpenAiUsageTests {
private DefaultUsage getDefaultUsage(OpenAiApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
@Test
void whenPromptTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getPromptTokens()).isEqualTo(200);
}
@Test
void whenPromptTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, null, 100);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getPromptTokens()).isEqualTo(0);
}
@Test
void whenGenerationTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getGenerationTokens()).isEqualTo(100);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getCompletionTokens()).isEqualTo(100);
}
@Test
void whenGenerationTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(null, 200, 200);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getGenerationTokens()).isEqualTo(0);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getCompletionTokens()).isEqualTo(0);
}
@Test
void whenTotalTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getTotalTokens()).isEqualTo(300);
}
@Test
void whenTotalTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, null);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getTotalTokens()).isEqualTo(300);
}
@Test
void whenPromptAndCompletionTokensDetailsIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, null);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getTotalTokens()).isEqualTo(300);
assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.promptTokensDetails()).isNull();
assertThat(nativeUsage.completionTokenDetails()).isNull();
}
@Test
void whenCompletionTokenDetailsIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, null);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getTotalTokens()).isEqualTo(300);
assertThat(usage.getReasoningTokens()).isEqualTo(0);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails()).isNull();
}
@Test
void whenReasoningTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, null));
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getReasoningTokens()).isEqualTo(0);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
}
@Test
void whenCompletionTokenDetailsIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(50, null, null, null));
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(50);
assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(50);
assertThat(nativeUsage.completionTokenDetails().acceptedPredictionTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().audioTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().rejectedPredictionTokens()).isEqualTo(null);
}
@Test
void whenAcceptedPredictionTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, 75, null, null));
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(75);
assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().acceptedPredictionTokens()).isEqualTo(75);
assertThat(nativeUsage.completionTokenDetails().audioTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().rejectedPredictionTokens()).isEqualTo(null);
}
@Test
void whenAudioTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, null, 125, null));
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(125);
assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().acceptedPredictionTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().audioTokens()).isEqualTo(125);
assertThat(nativeUsage.completionTokenDetails().rejectedPredictionTokens()).isEqualTo(null);
}
@Test
void whenRejectedPredictionTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, null));
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().acceptedPredictionTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().audioTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().rejectedPredictionTokens()).isEqualTo(null);
assertThat(nativeUsage.promptTokensDetails()).isEqualTo(null);
}
@@ -148,29 +161,32 @@ class OpenAiUsageTests {
void whenRejectedPredictionTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, 25));
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(25);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().acceptedPredictionTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().audioTokens()).isEqualTo(null);
assertThat(nativeUsage.completionTokenDetails().rejectedPredictionTokens()).isEqualTo(25);
}
@Test
void whenCacheTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300,
new OpenAiApi.Usage.PromptTokensDetails(null, null), null);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getPromptTokensDetails().audioTokens()).isEqualTo(0);
assertThat(usage.getPromptTokensDetails().cachedTokens()).isEqualTo(0);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.promptTokensDetails().audioTokens()).isEqualTo(null);
assertThat(nativeUsage.promptTokensDetails().cachedTokens()).isEqualTo(null);
}
@Test
void whenCacheTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300,
new OpenAiApi.Usage.PromptTokensDetails(99, 15), null);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getPromptTokensDetails().audioTokens()).isEqualTo(99);
assertThat(usage.getPromptTokensDetails().cachedTokens()).isEqualTo(15);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.promptTokensDetails().audioTokens()).isEqualTo(99);
assertThat(nativeUsage.promptTokensDetails().cachedTokens()).isEqualTo(15);
}
}

View File

@@ -30,6 +30,7 @@ import reactor.core.publisher.Mono;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
@@ -50,7 +51,6 @@ import org.springframework.ai.qianfan.api.QianFanApi.ChatCompletionMessage;
import org.springframework.ai.qianfan.api.QianFanApi.ChatCompletionMessage.Role;
import org.springframework.ai.qianfan.api.QianFanApi.ChatCompletionRequest;
import org.springframework.ai.qianfan.api.QianFanConstants;
import org.springframework.ai.qianfan.metadata.QianFanUsage;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
@@ -292,12 +292,16 @@ public class QianFanChatModel implements ChatModel, StreamingChatModel {
Assert.notNull(result, "QianFan ChatCompletionResult must not be null");
return ChatResponseMetadata.builder()
.id(result.id() != null ? result.id() : "")
.usage(result.usage() != null ? QianFanUsage.from(result.usage()) : new EmptyUsage())
.usage(result.usage() != null ? getDefaultUsage(result.usage()) : new EmptyUsage())
.model(model)
.keyValue("created", result.created() != null ? result.created() : 0L)
.build();
}
private DefaultUsage getDefaultUsage(QianFanApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
public void setObservationConvention(ChatModelObservationConvention observationConvention) {
this.observationConvention = observationConvention;
}

View File

@@ -22,6 +22,7 @@ import io.micrometer.observation.ObservationRegistry;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.document.Document;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.embedding.AbstractEmbeddingModel;
@@ -38,7 +39,6 @@ import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.qianfan.api.QianFanApi;
import org.springframework.ai.qianfan.api.QianFanApi.EmbeddingList;
import org.springframework.ai.qianfan.api.QianFanConstants;
import org.springframework.ai.qianfan.metadata.QianFanUsage;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.lang.Nullable;
import org.springframework.retry.support.RetryTemplate;
@@ -176,7 +176,7 @@ public class QianFanEmbeddingModel extends AbstractEmbeddingModel {
}
var metadata = new EmbeddingResponseMetadata(apiRequest.model(),
QianFanUsage.from(apiEmbeddingResponse.usage()));
getDefaultUsage(apiEmbeddingResponse.usage()));
List<Embedding> embeddings = apiEmbeddingResponse.data()
.stream()
@@ -192,6 +192,10 @@ public class QianFanEmbeddingModel extends AbstractEmbeddingModel {
}
private DefaultUsage getDefaultUsage(QianFanApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
/**
* Merge runtime and default {@link EmbeddingOptions} to compute the final options to
* use in the request.

View File

@@ -1,65 +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.qianfan.metadata;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.qianfan.api.QianFanApi;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal QianFan}.
*
* @author Thomas Vitale
*/
public class QianFanUsage implements Usage {
private final QianFanApi.Usage usage;
protected QianFanUsage(QianFanApi.Usage usage) {
Assert.notNull(usage, "QianFan Usage must not be null");
this.usage = usage;
}
public static QianFanUsage from(QianFanApi.Usage usage) {
return new QianFanUsage(usage);
}
protected QianFanApi.Usage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
return getUsage().promptTokens().longValue();
}
@Override
public Long getGenerationTokens() {
return 0L;
}
@Override
public Long getTotalTokens() {
return getUsage().totalTokens().longValue();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -148,7 +148,7 @@ public class QianFanChatModelObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

View File

@@ -1,44 +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.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);
}
}

View File

@@ -32,6 +32,7 @@ import com.google.protobuf.Value;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.DocumentEmbeddingModel;
@@ -44,7 +45,6 @@ import org.springframework.ai.embedding.EmbeddingResultMetadata.ModalityType;
import org.springframework.ai.model.Media;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingConnectionDetails;
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;
@@ -242,10 +242,14 @@ public class VertexAiMultimodalEmbeddingModel implements DocumentEmbeddingModel
private EmbeddingResponseMetadata generateResponseMetadata(String model, Integer totalTokens,
Map<String, Object> metadataToUse) {
Usage usage = new VertexAiEmbeddingUsage(totalTokens);
Usage usage = getDefaultUsage(totalTokens);
return new EmbeddingResponseMetadata(model, usage, metadataToUse);
}
private DefaultUsage getDefaultUsage(Integer totalTokens) {
return new DefaultUsage(0, 0, totalTokens);
}
@Override
public int dimensions() {
return KNOWN_EMBEDDING_DIMENSIONS.getOrDefault(this.defaultOptions.getModel(), 768);

View File

@@ -30,6 +30,7 @@ import com.google.cloud.aiplatform.v1.PredictionServiceClient;
import com.google.protobuf.Value;
import io.micrometer.observation.ObservationRegistry;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.AbstractEmbeddingModel;
@@ -45,7 +46,6 @@ import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.observation.conventions.AiProvider;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingConnectionDetails;
import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUsage;
import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils;
import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils.TextInstanceBuilder;
import org.springframework.ai.vertexai.embedding.VertexAiEmbeddingUtils.TextParametersBuilder;
@@ -222,11 +222,15 @@ public class VertexAiTextEmbeddingModel extends AbstractEmbeddingModel {
private EmbeddingResponseMetadata generateResponseMetadata(String model, Integer totalTokens) {
EmbeddingResponseMetadata metadata = new EmbeddingResponseMetadata();
metadata.setModel(model);
Usage usage = new VertexAiEmbeddingUsage(totalTokens);
Usage usage = getDefaultUsage(totalTokens);
metadata.setUsage(usage);
return metadata;
}
private DefaultUsage getDefaultUsage(Integer totalTokens) {
return new DefaultUsage(0, 0, totalTokens);
}
@Override
public int dimensions() {
return KNOWN_EMBEDDING_DIMENSIONS.getOrDefault(this.defaultOptions.getModel(), super.dimensions());

View File

@@ -57,6 +57,7 @@ import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.model.AbstractToolCallSupport;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
@@ -77,7 +78,6 @@ import org.springframework.ai.model.function.FunctionCallingOptions;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.ai.vertexai.gemini.common.VertexAiGeminiConstants;
import org.springframework.ai.vertexai.gemini.common.VertexAiGeminiSafetySetting;
import org.springframework.ai.vertexai.gemini.metadata.VertexAiUsage;
import org.springframework.beans.factory.DisposableBean;
import org.springframework.lang.NonNull;
import org.springframework.retry.support.RetryTemplate;
@@ -428,7 +428,12 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
}
private ChatResponseMetadata toChatResponseMetadata(GenerateContentResponse response) {
return ChatResponseMetadata.builder().usage(new VertexAiUsage(response.getUsageMetadata())).build();
return ChatResponseMetadata.builder().usage(getDefaultUsage(response.getUsageMetadata())).build();
}
private DefaultUsage getDefaultUsage(GenerateContentResponse.UsageMetadata usageMetadata) {
return new DefaultUsage(usageMetadata.getPromptTokenCount(), usageMetadata.getCandidatesTokenCount(),
usageMetadata.getTotalTokenCount(), usageMetadata);
}
private VertexAiGeminiChatOptions vertexAiGeminiChatOptions(Prompt prompt) {

View File

@@ -1,50 +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.vertexai.gemini.metadata;
import com.google.cloud.vertexai.api.GenerateContentResponse.UsageMetadata;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**
* Represents the usage of a Vertex AI model.
*
* @author Christian Tzolov
* @since 0.8.1
*
*/
public class VertexAiUsage implements Usage {
private final UsageMetadata usageMetadata;
public VertexAiUsage(UsageMetadata usageMetadata) {
Assert.notNull(usageMetadata, "UsageMetadata must not be null");
this.usageMetadata = usageMetadata;
}
@Override
public Long getPromptTokens() {
return Long.valueOf(this.usageMetadata.getPromptTokenCount());
}
@Override
public Long getGenerationTokens() {
return Long.valueOf(this.usageMetadata.getCandidatesTokenCount());
}
}

View File

@@ -148,7 +148,7 @@ public class VertexAiChatModelObservationIT {
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(
ChatModelObservationDocumentation.HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(
ChatModelObservationDocumentation.HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))

View File

@@ -38,6 +38,7 @@ import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.model.AbstractToolCallSupport;
import org.springframework.ai.chat.model.ChatModel;
@@ -68,7 +69,6 @@ import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionMessage.Role;
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionMessage.ToolCall;
import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionRequest;
import org.springframework.ai.zhipuai.api.ZhiPuApiConstants;
import org.springframework.ai.zhipuai.metadata.ZhiPuAiUsage;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
@@ -326,13 +326,17 @@ public class ZhiPuAiChatModel extends AbstractToolCallSupport implements ChatMod
Assert.notNull(result, "ZhiPuAI ChatCompletionResult must not be null");
return ChatResponseMetadata.builder()
.id(result.id() != null ? result.id() : "")
.usage(result.usage() != null ? ZhiPuAiUsage.from(result.usage()) : new EmptyUsage())
.usage(result.usage() != null ? getDefaultUsage(result.usage()) : new EmptyUsage())
.model(result.model() != null ? result.model() : "")
.keyValue("created", result.created() != null ? result.created() : 0L)
.keyValue("system-fingerprint", result.systemFingerprint() != null ? result.systemFingerprint() : "")
.build();
}
private DefaultUsage getDefaultUsage(ZhiPuAiApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
/**
* Convert the ChatCompletionChunk into a ChatCompletion. The Usage is set to null.
* @param chunk the ChatCompletionChunk to convert

View File

@@ -24,6 +24,7 @@ import io.micrometer.observation.ObservationRegistry;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.document.Document;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.embedding.AbstractEmbeddingModel;
@@ -40,7 +41,6 @@ import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.ai.zhipuai.api.ZhiPuAiApi;
import org.springframework.ai.zhipuai.api.ZhiPuApiConstants;
import org.springframework.ai.zhipuai.metadata.ZhiPuAiUsage;
import org.springframework.lang.Nullable;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
@@ -190,7 +190,7 @@ public class ZhiPuAiEmbeddingModel extends AbstractEmbeddingModel {
String model = (request.getOptions() != null && request.getOptions().getModel() != null)
? request.getOptions().getModel() : "unknown";
var metadata = new EmbeddingResponseMetadata(model, ZhiPuAiUsage.from(totalUsage));
var metadata = new EmbeddingResponseMetadata(model, getDefaultUsage(totalUsage));
var indexCounter = new AtomicInteger(0);
@@ -206,6 +206,10 @@ public class ZhiPuAiEmbeddingModel extends AbstractEmbeddingModel {
});
}
private DefaultUsage getDefaultUsage(ZhiPuAiApi.Usage usage) {
return new DefaultUsage(usage.promptTokens(), usage.completionTokens(), usage.totalTokens(), usage);
}
/**
* Merge runtime and default {@link EmbeddingOptions} to compute the final options to
* use in the request.

View File

@@ -1,66 +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.zhipuai.metadata;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.zhipuai.api.ZhiPuAiApi;
import org.springframework.util.Assert;
/**
* {@link Usage} implementation for {@literal ZhiPuAI}.
*
* @author Geng Rong
* @since 1.0.0 M1
*/
public class ZhiPuAiUsage implements Usage {
private final ZhiPuAiApi.Usage usage;
protected ZhiPuAiUsage(ZhiPuAiApi.Usage usage) {
Assert.notNull(usage, "ZhiPuAI Usage must not be null");
this.usage = usage;
}
public static ZhiPuAiUsage from(ZhiPuAiApi.Usage usage) {
return new ZhiPuAiUsage(usage);
}
protected ZhiPuAiApi.Usage getUsage() {
return this.usage;
}
@Override
public Long getPromptTokens() {
return getUsage().promptTokens().longValue();
}
@Override
public Long getGenerationTokens() {
return getUsage().completionTokens().longValue();
}
@Override
public Long getTotalTokens() {
return getUsage().totalTokens().longValue();
}
@Override
public String toString() {
return getUsage().toString();
}
}

View File

@@ -143,7 +143,7 @@ public class ZhiPuAiChatModelObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_INPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getPromptTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getGenerationTokens()))
String.valueOf(responseMetadata.getUsage().getCompletionTokens()))
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.USAGE_TOTAL_TOKENS.asString(),
String.valueOf(responseMetadata.getUsage().getTotalTokens()))
.hasBeenStarted()

22
pom.xml
View File

@@ -879,6 +879,28 @@
<skip>true</skip>
</configuration>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-checkstyle-plugin</artifactId>
<version>${maven-checkstyle-plugin.version}</version>
<reportSets>
<reportSet>
<reports>
<report>checkstyle</report>
</reports>
</reportSet>
</reportSets>
<configuration>
<configLocation>src/checkstyle/checkstyle.xml</configLocation>
<headerLocation>src/checkstyle/checkstyle-header.txt</headerLocation>
<propertyExpansion>
checkstyle.build.directory=${project.build.directory}
checkstyle.suppressions.file=${project.basedir}/src/checkstyle/checkstyle-suppressions.xml
checkstyle.additional.suppressions.file=${project.basedir}/src/checkstyle/checkstyle-suppressions.xml
checkstyle.header.file=${project.basedir}/src/checkstyle/checkstyle-header.txt
</propertyExpansion>
</configuration>
</plugin>
</plugins>
</reporting>

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2023-2024 the original author or authors.
* Copyright 2023-2025 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.
@@ -19,71 +19,127 @@ package org.springframework.ai.chat.metadata;
import java.util.Objects;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.annotation.JsonPropertyOrder;
/**
* Default implementation of the {@link Usage} interface.
*
* @author Mark Pollack
* @author Ilayaperumal Gopinathan
* @since 1.0.0
*/
@JsonPropertyOrder({ "promptTokens", "completionTokens", "totalTokens", "generationTokens", "nativeUsage" })
public class DefaultUsage implements Usage {
private final Long promptTokens;
private final Integer promptTokens;
private final Long generationTokens;
private final Long totalTokens;
private final Integer completionTokens;
/**
* Create a new DefaultUsage with promptTokens and generationTokens.
* @deprecated as of 1.0.0-M6, scheduled for removal
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
private final Long generationTokens;
private final int totalTokens;
private final Object nativeUsage;
/**
* Create a new DefaultUsage with promptTokens, completionTokens, totalTokens and
* native {@link Usage} object.
* @param promptTokens the number of tokens in the prompt, or {@code null} if not
* available
* @param generationTokens the number of tokens in the generation, or {@code null} if
* @param completionTokens the number of tokens in the generation, or {@code null} if
* not available
* @param totalTokens the total number of tokens, or {@code null} to calculate from
* promptTokens and completionTokens
* @param nativeUsage the native usage object returned by the model provider, or
* {@code null} to return the map of prompt, completion and total tokens.
*/
public DefaultUsage(Long promptTokens, Long generationTokens) {
this(promptTokens, generationTokens, null);
public DefaultUsage(Integer promptTokens, Integer completionTokens, Integer totalTokens, Object nativeUsage) {
this.promptTokens = promptTokens != null ? promptTokens : 0;
this.completionTokens = completionTokens != null ? completionTokens : 0;
this.generationTokens = Long.valueOf(this.completionTokens);
this.totalTokens = totalTokens != null ? totalTokens
: calculateTotalTokens(this.promptTokens, this.completionTokens);
this.nativeUsage = nativeUsage;
}
/**
* Create a new DefaultUsage with promptTokens, generationTokens, and totalTokens.
* Create a new DefaultUsage with promptTokens and completionTokens.
* @param promptTokens the number of tokens in the prompt, or {@code null} if not
* available
* @param generationTokens the number of tokens in the generation, or {@code null} if
* @param completionTokens the number of tokens in the generation, or {@code null} if
* not available
*/
public DefaultUsage(Integer promptTokens, Integer completionTokens) {
this(promptTokens, completionTokens, null, null);
}
/**
* Create a new DefaultUsage with promptTokens, completionTokens, and totalTokens.
* @param promptTokens the number of tokens in the prompt, or {@code null} if not
* available
* @param completionTokens the number of tokens in the generation, or {@code null} if
* not available
* @param totalTokens the total number of tokens, or {@code null} to calculate from
* promptTokens and generationTokens
* promptTokens and completionTokens
*/
public DefaultUsage(Integer promptTokens, Integer completionTokens, Integer totalTokens) {
this(promptTokens, completionTokens, totalTokens, null);
}
/**
* Create a new DefaultUsage with promptTokens, completionTokens, and totalTokens.
* This constructor is used for JSON deserialization and handles both the new format
* with completionTokens and the legacy format with generationTokens.
* @param promptTokens the number of tokens in the prompt
* @param completionTokens the number of tokens in the completion (new format)
* @param generationTokens the number of tokens in the generation (legacy format)
* @param totalTokens the total number of tokens
* @param nativeUsage the native usage object
* @return a new DefaultUsage instance
*/
@JsonCreator
public DefaultUsage(@JsonProperty("promptTokens") Long promptTokens,
@JsonProperty("generationTokens") Long generationTokens, @JsonProperty("totalTokens") Long totalTokens) {
this.promptTokens = promptTokens != null ? promptTokens : 0L;
this.generationTokens = generationTokens != null ? generationTokens : 0L;
this.totalTokens = totalTokens != null ? totalTokens
: calculateTotalTokens(this.promptTokens, this.generationTokens);
public static DefaultUsage fromJson(@JsonProperty("promptTokens") Integer promptTokens,
@JsonProperty("completionTokens") Integer completionTokens,
@JsonProperty("generationTokens") Long generationTokens, @JsonProperty("totalTokens") Integer totalTokens,
@JsonProperty("nativeUsage") Object nativeUsage) {
Integer effectiveCompletionTokens = completionTokens != null ? completionTokens
: (generationTokens != null ? generationTokens.intValue() : 0);
return new DefaultUsage(promptTokens, effectiveCompletionTokens, totalTokens, nativeUsage);
}
@Override
@JsonProperty("promptTokens")
public Long getPromptTokens() {
public Integer getPromptTokens() {
return this.promptTokens;
}
@Override
@JsonProperty("generationTokens")
public Long getGenerationTokens() {
return this.generationTokens;
@JsonProperty("completionTokens")
public Integer getCompletionTokens() {
return this.completionTokens;
}
@Override
@JsonProperty("totalTokens")
public Long getTotalTokens() {
public Integer getTotalTokens() {
return this.totalTokens;
}
private Long calculateTotalTokens(Long promptTokens, Long generationTokens) {
return promptTokens + generationTokens;
@Override
@JsonProperty("nativeUsage")
@JsonInclude(JsonInclude.Include.NON_NULL)
public Object getNativeUsage() {
return this.nativeUsage;
}
private Integer calculateTotalTokens(Integer promptTokens, Integer completionTokens) {
return promptTokens + completionTokens;
}
@Override
@@ -94,20 +150,25 @@ public class DefaultUsage implements Usage {
if (o == null || getClass() != o.getClass()) {
return false;
}
DefaultUsage that = (DefaultUsage) o;
return Objects.equals(this.promptTokens, that.promptTokens)
&& Objects.equals(this.generationTokens, that.generationTokens)
&& Objects.equals(this.totalTokens, that.totalTokens);
return this.totalTokens == that.totalTokens && Objects.equals(this.promptTokens, that.promptTokens)
&& Objects.equals(this.completionTokens, that.completionTokens)
&& Objects.equals(this.nativeUsage, that.nativeUsage);
}
@Override
public int hashCode() {
return Objects.hash(this.promptTokens, this.generationTokens, this.totalTokens);
int result = Objects.hashCode(this.promptTokens);
result = 31 * result + Objects.hashCode(this.completionTokens);
result = 31 * result + this.totalTokens;
result = 31 * result + Objects.hashCode(this.nativeUsage);
return result;
}
@Override
public String toString() {
return "DefaultUsage{" + "promptTokens=" + this.promptTokens + ", generationTokens=" + this.generationTokens
return "DefaultUsage{" + "promptTokens=" + this.promptTokens + ", completionTokens=" + this.completionTokens
+ ", totalTokens=" + this.totalTokens + '}';
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2023-2024 the original author or authors.
* Copyright 2023-2025 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.
@@ -16,22 +16,30 @@
package org.springframework.ai.chat.metadata;
import java.util.Map;
/**
* A EmpytUsage implementation that returns zero for all property getters
*
* @author John Blum
* @author Ilayaperumal Gopinathan
* @since 0.7.0
*/
public class EmptyUsage implements Usage {
@Override
public Long getPromptTokens() {
return 0L;
public Integer getPromptTokens() {
return 0;
}
@Override
public Long getGenerationTokens() {
return 0L;
public Integer getCompletionTokens() {
return 0;
}
@Override
public Object getNativeUsage() {
return Map.of();
}
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2023-2024 the original author or authors.
* Copyright 2023-2025 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.
@@ -21,26 +21,32 @@ package org.springframework.ai.chat.metadata;
* per AI request.
*
* @author John Blum
* @author Ilayaperumal Gopinathan
* @since 0.7.0
*/
public interface Usage {
/**
* Returns the number of tokens used in the {@literal prompt} of the AI request.
* @return an {@link Long} with the number of tokens used in the {@literal prompt} of
* the AI request.
* @see #getGenerationTokens()
* @return an {@link Integer} with the number of tokens used in the {@literal prompt}
* of the AI request.
* @see #getCompletionTokens()
*/
Long getPromptTokens();
Integer getPromptTokens();
@Deprecated(forRemoval = true, since = "1.0.0-M6")
default Long getGenerationTokens() {
return getCompletionTokens().longValue();
}
/**
* Returns the number of tokens returned in the {@literal generation (aka completion)}
* of the AI's response.
* @return an {@link Long} with the number of tokens returned in the
* @return an {@link Integer} with the number of tokens returned in the
* {@literal generation (aka completion)} of the AI's response.
* @see #getPromptTokens()
*/
Long getGenerationTokens();
Integer getCompletionTokens();
/**
* Return the total number of tokens from both the {@literal prompt} of an AI request
@@ -48,14 +54,20 @@ public interface Usage {
* @return the total number of tokens from both the {@literal prompt} of an AI request
* and {@literal generation} of the AI's response.
* @see #getPromptTokens()
* @see #getGenerationTokens()
* @see #getCompletionTokens()
*/
default Long getTotalTokens() {
Long promptTokens = getPromptTokens();
default Integer getTotalTokens() {
Integer promptTokens = getPromptTokens();
promptTokens = promptTokens != null ? promptTokens : 0;
Long completionTokens = getGenerationTokens();
Integer completionTokens = getCompletionTokens();
completionTokens = completionTokens != null ? completionTokens : 0;
return promptTokens + completionTokens;
}
/**
* Return the usage data from the underlying model API response.
* @return the object of type inferred by the API response.
*/
Object getNativeUsage();
}

View File

@@ -50,12 +50,12 @@ public final class UsageUtils {
// For a valid usage from previous chat response, accumulate it to the current
// usage.
if (!isEmpty(currentUsage)) {
Long promptTokens = currentUsage.getPromptTokens().longValue();
Long generationTokens = currentUsage.getGenerationTokens().longValue();
Long totalTokens = currentUsage.getTotalTokens().longValue();
Integer promptTokens = currentUsage.getPromptTokens();
Integer generationTokens = currentUsage.getCompletionTokens();
Integer totalTokens = currentUsage.getTotalTokens();
// Make sure to accumulate the usage from the previous chat response.
promptTokens += usageFromPreviousChatResponse.getPromptTokens();
generationTokens += usageFromPreviousChatResponse.getGenerationTokens();
generationTokens += usageFromPreviousChatResponse.getCompletionTokens();
totalTokens += usageFromPreviousChatResponse.getTotalTokens();
return new DefaultUsage(promptTokens, generationTokens, totalTokens);
}

View File

@@ -81,9 +81,9 @@ public class MessageAggregator {
ChatGenerationMetadata.NULL);
// Usage
AtomicReference<Long> metadataUsagePromptTokensRef = new AtomicReference<>(0L);
AtomicReference<Long> metadataUsageGenerationTokensRef = new AtomicReference<>(0L);
AtomicReference<Long> metadataUsageTotalTokensRef = new AtomicReference<>(0L);
AtomicReference<Integer> metadataUsagePromptTokensRef = new AtomicReference<Integer>(0);
AtomicReference<Integer> metadataUsageGenerationTokensRef = new AtomicReference<Integer>(0);
AtomicReference<Integer> metadataUsageTotalTokensRef = new AtomicReference<Integer>(0);
AtomicReference<PromptMetadata> metadataPromptMetadataRef = new AtomicReference<>(PromptMetadata.empty());
AtomicReference<RateLimit> metadataRateLimitRef = new AtomicReference<>(new EmptyRateLimit());
@@ -96,9 +96,9 @@ public class MessageAggregator {
messageMetadataMapRef.set(new HashMap<>());
metadataIdRef.set("");
metadataModelRef.set("");
metadataUsagePromptTokensRef.set(0L);
metadataUsageGenerationTokensRef.set(0L);
metadataUsageTotalTokensRef.set(0L);
metadataUsagePromptTokensRef.set(0);
metadataUsageGenerationTokensRef.set(0);
metadataUsageTotalTokensRef.set(0);
metadataPromptMetadataRef.set(PromptMetadata.empty());
metadataRateLimitRef.set(new EmptyRateLimit());
@@ -121,7 +121,7 @@ public class MessageAggregator {
Usage usage = chatResponse.getMetadata().getUsage();
metadataUsagePromptTokensRef.set(
usage.getPromptTokens() > 0 ? usage.getPromptTokens() : metadataUsagePromptTokensRef.get());
metadataUsageGenerationTokensRef.set(usage.getGenerationTokens() > 0 ? usage.getGenerationTokens()
metadataUsageGenerationTokensRef.set(usage.getCompletionTokens() > 0 ? usage.getCompletionTokens()
: metadataUsageGenerationTokensRef.get());
metadataUsageTotalTokensRef
.set(usage.getTotalTokens() > 0 ? usage.getTotalTokens() : metadataUsageTotalTokensRef.get());
@@ -162,32 +162,40 @@ public class MessageAggregator {
messageMetadataMapRef.set(new HashMap<>());
metadataIdRef.set("");
metadataModelRef.set("");
metadataUsagePromptTokensRef.set(0L);
metadataUsageGenerationTokensRef.set(0L);
metadataUsageTotalTokensRef.set(0L);
metadataUsagePromptTokensRef.set(0);
metadataUsageGenerationTokensRef.set(0);
metadataUsageTotalTokensRef.set(0);
metadataPromptMetadataRef.set(PromptMetadata.empty());
metadataRateLimitRef.set(new EmptyRateLimit());
}).doOnError(e -> logger.error("Aggregation Error", e));
}
public record DefaultUsage(long promptTokens, long generationTokens, long totalTokens) implements Usage {
public record DefaultUsage(Integer promptTokens, Integer completionTokens, Integer totalTokens) implements Usage {
@Override
public Long getPromptTokens() {
public Integer getPromptTokens() {
return promptTokens();
}
@Override
public Long getGenerationTokens() {
return generationTokens();
public Integer getCompletionTokens() {
return completionTokens();
}
@Override
public Long getTotalTokens() {
public Integer getTotalTokens() {
return totalTokens();
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", promptTokens());
usage.put("completionTokens", completionTokens());
usage.put("totalTokens", totalTokens());
return usage;
}
}
}

View File

@@ -222,10 +222,10 @@ public class DefaultChatModelObservationConvention implements ChatModelObservati
protected KeyValues usageOutputTokens(KeyValues keyValues, ChatModelObservationContext context) {
if (context.getResponse() != null && context.getResponse().getMetadata() != null
&& context.getResponse().getMetadata().getUsage() != null
&& context.getResponse().getMetadata().getUsage().getGenerationTokens() != null) {
&& context.getResponse().getMetadata().getUsage().getCompletionTokens() != null) {
return keyValues.and(
ChatModelObservationDocumentation.HighCardinalityKeyNames.USAGE_OUTPUT_TOKENS.asString(),
String.valueOf(context.getResponse().getMetadata().getUsage().getGenerationTokens()));
String.valueOf(context.getResponse().getMetadata().getUsage().getCompletionTokens()));
}
return keyValues;
}

View File

@@ -54,13 +54,13 @@ public final class ModelUsageMetricsGenerator {
.increment(usage.getPromptTokens());
}
if (usage.getGenerationTokens() != null) {
if (usage.getCompletionTokens() != null) {
Counter.builder(AiObservationMetricNames.TOKEN_USAGE.value())
.tag(AiObservationMetricAttributes.TOKEN_TYPE.value(), AiTokenType.OUTPUT.value())
.description(DESCRIPTION)
.tags(createTags(context))
.register(meterRegistry)
.increment(usage.getGenerationTokens());
.increment(usage.getCompletionTokens());
}
if (usage.getTotalTokens() != null) {

View File

@@ -104,7 +104,7 @@ public class QuestionAnswerAdvisorTests {
public Duration getTokensReset() {
return Duration.ofSeconds(9);
}
}).usage(new DefaultUsage(6L, 7L))
}).usage(new DefaultUsage(6, 7))
.build()));
// @formatter:on
@@ -137,7 +137,7 @@ public class QuestionAnswerAdvisorTests {
assertThat(response.getMetadata().getRateLimit().getTokensRemaining()).isEqualTo(8L);
assertThat(response.getMetadata().getRateLimit().getTokensReset()).isEqualTo(Duration.ofSeconds(9));
assertThat(response.getMetadata().getUsage().getPromptTokens()).isEqualTo(6L);
assertThat(response.getMetadata().getUsage().getGenerationTokens()).isEqualTo(7L);
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isEqualTo(7L);
assertThat(response.getMetadata().getUsage().getTotalTokens()).isEqualTo(6L + 7L);
assertThat(response.getMetadata().get("key6").toString()).isEqualTo("value6");
assertThat(response.getMetadata().get("key1").toString()).isEqualTo("value1");

View File

@@ -19,7 +19,10 @@ package org.springframework.ai.chat.metadata;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.junit.jupiter.api.Test;
import static org.junit.jupiter.api.Assertions.assertEquals;
import java.util.HashMap;
import java.util.Map;
import static org.assertj.core.api.Assertions.assertThat;
public class DefaultUsageTests {
@@ -27,93 +30,255 @@ public class DefaultUsageTests {
@Test
void testSerializationWithAllFields() throws Exception {
DefaultUsage usage = new DefaultUsage(100L, 50L, 150L);
DefaultUsage usage = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150));
String json = this.objectMapper.writeValueAsString(usage);
assertEquals("{\"promptTokens\":100,\"generationTokens\":50,\"totalTokens\":150}", json);
assertThat(json)
.isEqualTo("{\"promptTokens\":100,\"completionTokens\":50,\"totalTokens\":150,\"generationTokens\":50}");
}
@Test
void testDeserializationWithAllFields() throws Exception {
String json = "{\"promptTokens\":100,\"generationTokens\":50,\"totalTokens\":150}";
String json = "{\"promptTokens\":100,\"completionTokens\":50,\"totalTokens\":150,\"generationTokens\":50}";
DefaultUsage usage = this.objectMapper.readValue(json, DefaultUsage.class);
assertEquals(100L, usage.getPromptTokens());
assertEquals(50L, usage.getGenerationTokens());
assertEquals(150L, usage.getTotalTokens());
assertThat(usage.getPromptTokens()).isEqualTo(100);
assertThat(usage.getCompletionTokens()).isEqualTo(50);
assertThat(usage.getTotalTokens()).isEqualTo(150);
}
@Test
void testSerializationWithNullFields() throws Exception {
DefaultUsage usage = new DefaultUsage(null, null, null);
DefaultUsage usage = new DefaultUsage((Integer) null, (Integer) null, (Integer) null);
String json = this.objectMapper.writeValueAsString(usage);
assertEquals("{\"promptTokens\":0,\"generationTokens\":0,\"totalTokens\":0}", json);
assertThat(json)
.isEqualTo("{\"promptTokens\":0,\"completionTokens\":0,\"totalTokens\":0,\"generationTokens\":0}");
}
@Test
void testDeserializationWithMissingFields() throws Exception {
String json = "{\"promptTokens\":100}";
DefaultUsage usage = this.objectMapper.readValue(json, DefaultUsage.class);
assertEquals(100L, usage.getPromptTokens());
assertEquals(0L, usage.getGenerationTokens());
assertEquals(100L, usage.getTotalTokens());
assertThat(usage.getPromptTokens()).isEqualTo(100);
assertThat(usage.getCompletionTokens()).isEqualTo(0);
assertThat(usage.getTotalTokens()).isEqualTo(100);
}
@Test
void testDeserializationWithNullFields() throws Exception {
String json = "{\"promptTokens\":null,\"generationTokens\":null,\"totalTokens\":null}";
String json = "{\"promptTokens\":null,\"completionTokens\":null,\"totalTokens\":null}";
DefaultUsage usage = this.objectMapper.readValue(json, DefaultUsage.class);
assertEquals(0L, usage.getPromptTokens());
assertEquals(0L, usage.getGenerationTokens());
assertEquals(0L, usage.getTotalTokens());
assertThat(usage.getPromptTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokens()).isEqualTo(0);
assertThat(usage.getTotalTokens()).isEqualTo(0);
}
@Test
void testRoundTripSerialization() throws Exception {
DefaultUsage original = new DefaultUsage(100L, 50L, 150L);
DefaultUsage original = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150));
String json = this.objectMapper.writeValueAsString(original);
DefaultUsage deserialized = this.objectMapper.readValue(json, DefaultUsage.class);
assertEquals(original.getPromptTokens(), deserialized.getPromptTokens());
assertEquals(original.getGenerationTokens(), deserialized.getGenerationTokens());
assertEquals(original.getTotalTokens(), deserialized.getTotalTokens());
assertThat(deserialized.getPromptTokens()).isEqualTo(original.getPromptTokens());
assertThat(deserialized.getCompletionTokens()).isEqualTo(original.getCompletionTokens());
assertThat(deserialized.getTotalTokens()).isEqualTo(original.getTotalTokens());
}
@Test
void testTwoArgumentConstructorAndSerialization() throws Exception {
DefaultUsage usage = new DefaultUsage(100L, 50L);
DefaultUsage usage = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50));
// Test that the fields are set correctly
assertEquals(100L, usage.getPromptTokens());
assertEquals(50L, usage.getGenerationTokens());
assertEquals(150L, usage.getTotalTokens()); // 100 + 50 = 150
assertThat(usage.getPromptTokens()).isEqualTo(100);
assertThat(usage.getCompletionTokens()).isEqualTo(50);
assertThat(usage.getTotalTokens()).isEqualTo(150); // 100 + 50 = 150
// Test serialization
String json = this.objectMapper.writeValueAsString(usage);
assertEquals("{\"promptTokens\":100,\"generationTokens\":50,\"totalTokens\":150}", json);
assertThat(json)
.isEqualTo("{\"promptTokens\":100,\"completionTokens\":50,\"totalTokens\":150,\"generationTokens\":50}");
// Test deserialization
DefaultUsage deserializedUsage = this.objectMapper.readValue(json, DefaultUsage.class);
assertEquals(100L, deserializedUsage.getPromptTokens());
assertEquals(50L, deserializedUsage.getGenerationTokens());
assertEquals(150L, deserializedUsage.getTotalTokens());
assertThat(deserializedUsage.getPromptTokens()).isEqualTo(100);
assertThat(deserializedUsage.getCompletionTokens()).isEqualTo(50);
assertThat(deserializedUsage.getTotalTokens()).isEqualTo(150);
}
@Test
void testTwoArgumentConstructorWithNullValues() throws Exception {
DefaultUsage usage = new DefaultUsage(null, null);
DefaultUsage usage = new DefaultUsage((Integer) null, (Integer) null);
// Test that null values are converted to 0
assertEquals(0L, usage.getPromptTokens());
assertEquals(0L, usage.getGenerationTokens());
assertEquals(0L, usage.getTotalTokens());
assertThat(usage.getPromptTokens()).isEqualTo(0);
assertThat(usage.getCompletionTokens()).isEqualTo(0);
assertThat(usage.getTotalTokens()).isEqualTo(0);
// Test serialization
String json = this.objectMapper.writeValueAsString(usage);
assertEquals("{\"promptTokens\":0,\"generationTokens\":0,\"totalTokens\":0}", json);
assertThat(json)
.isEqualTo("{\"promptTokens\":0,\"completionTokens\":0,\"totalTokens\":0,\"generationTokens\":0}");
// Test deserialization
DefaultUsage deserializedUsage = this.objectMapper.readValue(json, DefaultUsage.class);
assertEquals(0L, deserializedUsage.getPromptTokens());
assertEquals(0L, deserializedUsage.getGenerationTokens());
assertEquals(0L, deserializedUsage.getTotalTokens());
assertThat(deserializedUsage.getPromptTokens()).isEqualTo(0);
assertThat(deserializedUsage.getCompletionTokens()).isEqualTo(0);
assertThat(deserializedUsage.getTotalTokens()).isEqualTo(0);
}
@Test
void testDeserializationWithLegacyFormat() throws Exception {
String json = "{\"promptTokens\":100,\"generationTokens\":50,\"totalTokens\":150}";
DefaultUsage usage = this.objectMapper.readValue(json, DefaultUsage.class);
assertThat(usage.getPromptTokens()).isEqualTo(100);
assertThat(usage.getCompletionTokens()).isEqualTo(50);
assertThat(usage.getTotalTokens()).isEqualTo(150);
}
@Test
void testDeserializationWithDifferentPropertyOrder() throws Exception {
String json = "{\"totalTokens\":150,\"generationTokens\":50,\"completionTokens\":50,\"promptTokens\":100}";
DefaultUsage usage = this.objectMapper.readValue(json, DefaultUsage.class);
assertThat(usage.getPromptTokens()).isEqualTo(100);
assertThat(usage.getCompletionTokens()).isEqualTo(50);
assertThat(usage.getTotalTokens()).isEqualTo(150);
}
@Test
void testSerializationWithCustomNativeUsage() throws Exception {
Map<String, Object> customNativeUsage = new HashMap<>();
customNativeUsage.put("custom_field", "custom_value");
customNativeUsage.put("custom_number", 42);
DefaultUsage usage = new DefaultUsage(100, 50, 150, customNativeUsage);
String json = this.objectMapper.writeValueAsString(usage);
assertThat(json).isEqualTo(
"{\"promptTokens\":100,\"completionTokens\":50,\"totalTokens\":150,\"generationTokens\":50,\"nativeUsage\":{\"custom_field\":\"custom_value\",\"custom_number\":42}}");
}
@Test
void testDeserializationWithCustomNativeUsage() throws Exception {
String json = "{\"promptTokens\":100,\"completionTokens\":50,\"totalTokens\":150,\"nativeUsage\":{\"custom_field\":\"custom_value\",\"custom_number\":42}}";
DefaultUsage usage = this.objectMapper.readValue(json, DefaultUsage.class);
assertThat(usage.getPromptTokens()).isEqualTo(100);
assertThat(usage.getCompletionTokens()).isEqualTo(50);
assertThat(usage.getTotalTokens()).isEqualTo(150);
@SuppressWarnings("unchecked")
Map<String, Object> nativeUsage = (Map<String, Object>) usage.getNativeUsage();
assertThat(nativeUsage.get("custom_field")).isEqualTo("custom_value");
assertThat(nativeUsage.get("custom_number")).isEqualTo(42);
}
@Test
void testArbitraryNativeUsageMap() throws Exception {
Map<String, Object> arbitraryMap = new HashMap<>();
arbitraryMap.put("field1", "value1");
arbitraryMap.put("field2", 42);
arbitraryMap.put("field3", true);
arbitraryMap.put("field4", java.util.Arrays.asList(1, 2, 3));
arbitraryMap.put("field5", java.util.Map.of("nested", "value"));
DefaultUsage usage = new DefaultUsage(100, 50, 150, arbitraryMap);
String json = this.objectMapper.writeValueAsString(usage);
DefaultUsage deserialized = this.objectMapper.readValue(json, DefaultUsage.class);
assertThat(deserialized.getPromptTokens()).isEqualTo(usage.getPromptTokens());
assertThat(deserialized.getCompletionTokens()).isEqualTo(usage.getCompletionTokens());
assertThat(deserialized.getTotalTokens()).isEqualTo(usage.getTotalTokens());
assertThat(deserialized.getGenerationTokens()).isEqualTo(usage.getGenerationTokens());
@SuppressWarnings("unchecked")
Map<String, Object> deserializedMap = (Map<String, Object>) deserialized.getNativeUsage();
assertThat(deserializedMap.get("field1")).isEqualTo("value1");
assertThat(deserializedMap.get("field2")).isEqualTo(42);
assertThat(deserializedMap.get("field3")).isEqualTo(true);
assertThat(deserializedMap.get("field4")).isEqualTo(java.util.Arrays.asList(1, 2, 3));
assertThat(deserializedMap.get("field5")).isEqualTo(java.util.Map.of("nested", "value"));
}
@Test
@SuppressWarnings("deprecation")
void testDeprecatedGenerationTokens() {
DefaultUsage usage = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150));
assertThat(usage.getGenerationTokens()).isEqualTo(50L);
assertThat(usage.getCompletionTokens().longValue()).isEqualTo(usage.getGenerationTokens());
}
@Test
void testEqualsAndHashCode() {
DefaultUsage usage1 = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150));
DefaultUsage usage2 = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150));
DefaultUsage usage3 = new DefaultUsage(Integer.valueOf(200), Integer.valueOf(100), Integer.valueOf(300));
DefaultUsage usage4 = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150),
Map.of("custom", "value"));
// Test equals
assertThat(usage1).isEqualTo(usage2);
assertThat(usage1).isNotEqualTo(usage3);
assertThat(usage1).isNotEqualTo(usage4);
assertThat(usage1).isNotEqualTo(null);
assertThat(usage1).isNotEqualTo(new Object());
// Test hashCode
assertThat(usage1).hasSameHashCodeAs(usage2);
assertThat(usage1.hashCode()).isNotEqualTo(usage3.hashCode());
assertThat(usage1.hashCode()).isNotEqualTo(usage4.hashCode());
// Test reflexivity
assertThat(usage1).isEqualTo(usage1);
assertThat(usage1).hasSameHashCodeAs(usage1);
// Test symmetry
assertThat(usage1.equals(usage2)).isEqualTo(usage2.equals(usage1));
// Test with different nativeUsage
DefaultUsage usage5 = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150),
Map.of("key", "value"));
DefaultUsage usage6 = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150),
Map.of("key", "value"));
assertThat(usage5).isEqualTo(usage6);
assertThat(usage5).hasSameHashCodeAs(usage6);
}
@Test
void testToString() {
DefaultUsage usage = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150));
assertThat(usage).hasToString("DefaultUsage{promptTokens=100, completionTokens=50, totalTokens=150}");
// Test with custom nativeUsage
DefaultUsage usageWithNative = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), Integer.valueOf(150),
Map.of("custom", "value"));
assertThat(usageWithNative).hasToString("DefaultUsage{promptTokens=100, completionTokens=50, totalTokens=150}");
// Test with null values
DefaultUsage usageWithNulls = new DefaultUsage(null, null, null);
assertThat(usageWithNulls).hasToString("DefaultUsage{promptTokens=0, completionTokens=0, totalTokens=0}");
}
@Test
void testNegativeTokenValues() throws Exception {
DefaultUsage usage = new DefaultUsage(Integer.valueOf(-1), Integer.valueOf(-2), Integer.valueOf(-3));
assertThat(usage.getPromptTokens()).isEqualTo(-1);
assertThat(usage.getCompletionTokens()).isEqualTo(-2);
assertThat(usage.getTotalTokens()).isEqualTo(-3);
String json = this.objectMapper.writeValueAsString(usage);
assertThat(json)
.isEqualTo("{\"promptTokens\":-1,\"completionTokens\":-2,\"totalTokens\":-3,\"generationTokens\":-2}");
}
@Test
void testCalculatedTotalTokens() {
// Test when total tokens is null and should be calculated
DefaultUsage usage = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50), null);
assertThat(usage.getTotalTokens()).isEqualTo(150); // Should be sum of prompt and
// completion tokens
// Test that explicit total tokens takes precedence over calculated
DefaultUsage usageWithExplicitTotal = new DefaultUsage(Integer.valueOf(100), Integer.valueOf(50),
Integer.valueOf(200));
assertThat(usageWithExplicitTotal.getTotalTokens()).isEqualTo(200); // Should use
// explicit
// value
}
}

View File

@@ -16,7 +16,9 @@
package org.springframework.ai.chat.observation;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import io.micrometer.core.instrument.MeterRegistry;
import io.micrometer.core.instrument.simple.SimpleMeterRegistry;
@@ -106,13 +108,22 @@ class ChatModelMeterObservationHandlerTests {
static class TestUsage implements Usage {
@Override
public Long getPromptTokens() {
return 1000L;
public Integer getPromptTokens() {
return 1000;
}
@Override
public Long getGenerationTokens() {
return 500L;
public Integer getCompletionTokens() {
return 500;
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", getPromptTokens());
usage.put("completionTokens", getCompletionTokens());
usage.put("totalTokens", getTotalTokens());
return usage;
}
}

View File

@@ -16,7 +16,9 @@
package org.springframework.ai.chat.observation;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import io.micrometer.common.KeyValue;
import io.micrometer.observation.Observation;
@@ -30,6 +32,7 @@ import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.lang.Nullable;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.chat.observation.ChatModelObservationDocumentation.HighCardinalityKeyNames;
@@ -183,13 +186,22 @@ class DefaultChatModelObservationConventionTests {
static class TestUsage implements Usage {
@Override
public Long getPromptTokens() {
return 1000L;
public Integer getPromptTokens() {
return 1000;
}
@Override
public Long getGenerationTokens() {
return 500L;
public Integer getCompletionTokens() {
return 500;
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", getPromptTokens());
usage.put("completionTokens", getCompletionTokens());
usage.put("totalTokens", getTotalTokens());
return usage;
}
}

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.embedding.observation;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
@@ -134,13 +135,22 @@ class DefaultEmbeddingModelObservationConventionTests {
static class TestUsage implements Usage {
@Override
public Long getPromptTokens() {
return 1000L;
public Integer getPromptTokens() {
return 1000;
}
@Override
public Long getGenerationTokens() {
return 0L;
public Integer getCompletionTokens() {
return 0;
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", getPromptTokens());
usage.put("completionTokens", getCompletionTokens());
usage.put("totalTokens", getTotalTokens());
return usage;
}
}

View File

@@ -16,8 +16,10 @@
package org.springframework.ai.embedding.observation;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import io.micrometer.core.instrument.MeterRegistry;
import io.micrometer.core.instrument.simple.SimpleMeterRegistry;
@@ -104,18 +106,27 @@ class EmbeddingModelMeterObservationHandlerTests {
static class TestUsage implements Usage {
@Override
public Long getPromptTokens() {
return 1000L;
public Integer getPromptTokens() {
return 1000;
}
@Override
public Long getGenerationTokens() {
return 0L;
public Integer getCompletionTokens() {
return 0;
}
@Override
public Long getTotalTokens() {
return 1000L;
public Integer getTotalTokens() {
return 1000;
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", getPromptTokens());
usage.put("completionTokens", getCompletionTokens());
usage.put("totalTokens", getTotalTokens());
return usage;
}
}

View File

@@ -36,10 +36,10 @@ import static org.mockito.Mockito.verifyNoMoreInteractions;
*/
public class UsageTests {
private Usage mockUsage(Long promptTokens, Long generationTokens) {
private Usage mockUsage(Integer promptTokens, Integer generationTokens) {
Usage mockUsage = mock(Usage.class);
doReturn(promptTokens).when(mockUsage).getPromptTokens();
doReturn(generationTokens).when(mockUsage).getGenerationTokens();
doReturn(generationTokens).when(mockUsage).getCompletionTokens();
doCallRealMethod().when(mockUsage).getTotalTokens();
return mockUsage;
}
@@ -47,7 +47,7 @@ public class UsageTests {
private void verifyUsage(Usage usage) {
verify(usage, times(1)).getTotalTokens();
verify(usage, times(1)).getPromptTokens();
verify(usage, times(1)).getGenerationTokens();
verify(usage, times(1)).getCompletionTokens();
verifyNoMoreInteractions(usage);
}
@@ -63,27 +63,27 @@ public class UsageTests {
@Test
void totalTokensEqualsPromptTokens() {
Usage usage = mockUsage(10L, null);
Usage usage = mockUsage(10, null);
assertThat(usage.getTotalTokens()).isEqualTo(10L);
assertThat(usage.getTotalTokens()).isEqualTo(10);
verifyUsage(usage);
}
@Test
void totalTokensEqualsGenerationTokens() {
Usage usage = mockUsage(null, 15L);
Usage usage = mockUsage(null, 15);
assertThat(usage.getTotalTokens()).isEqualTo(15L);
assertThat(usage.getTotalTokens()).isEqualTo(15);
verifyUsage(usage);
}
@Test
void totalTokensEqualsPromptTokensPlusGenerationTokens() {
Usage usage = mockUsage(10L, 15L);
Usage usage = mockUsage(10, 15);
assertThat(usage.getTotalTokens()).isEqualTo(25L);
assertThat(usage.getTotalTokens()).isEqualTo(25);
verifyUsage(usage);
}

View File

@@ -16,6 +16,9 @@
package org.springframework.ai.model.observation;
import java.util.HashMap;
import java.util.Map;
import io.micrometer.common.KeyValue;
import io.micrometer.core.instrument.simple.SimpleMeterRegistry;
import io.micrometer.observation.Observation;
@@ -38,7 +41,7 @@ class ModelUsageMetricsGeneratorTests {
@Test
void whenTokenUsageThenMetrics() {
var meterRegistry = new SimpleMeterRegistry();
var usage = new TestUsage(1000L, 500L, 1500L);
var usage = new TestUsage(1000, 500, 1500);
ModelUsageMetricsGenerator.generate(usage, buildContext(), meterRegistry);
assertThat(meterRegistry.get(AiObservationMetricNames.TOKEN_USAGE.value()).meters()).hasSize(3);
@@ -59,7 +62,7 @@ class ModelUsageMetricsGeneratorTests {
@Test
void whenPartialTokenUsageThenMetrics() {
var meterRegistry = new SimpleMeterRegistry();
var usage = new TestUsage(1000L, null, 1000L);
var usage = new TestUsage(1000, null, 1000);
ModelUsageMetricsGenerator.generate(usage, buildContext(), meterRegistry);
assertThat(meterRegistry.get(AiObservationMetricNames.TOKEN_USAGE.value()).meters()).hasSize(2);
@@ -82,33 +85,42 @@ class ModelUsageMetricsGeneratorTests {
static class TestUsage implements Usage {
private final Long promptTokens;
private final Integer promptTokens;
private final Long generationTokens;
private final Integer generationTokens;
private final Long totalTokens;
private final int totalTokens;
TestUsage(Long promptTokens, Long generationTokens, Long totalTokens) {
TestUsage(Integer promptTokens, Integer generationTokens, int totalTokens) {
this.promptTokens = promptTokens;
this.generationTokens = generationTokens;
this.totalTokens = totalTokens;
}
@Override
public Long getPromptTokens() {
public Integer getPromptTokens() {
return this.promptTokens;
}
@Override
public Long getGenerationTokens() {
public Integer getCompletionTokens() {
return this.generationTokens;
}
@Override
public Long getTotalTokens() {
public Integer getTotalTokens() {
return this.totalTokens;
}
@Override
public Map<String, Integer> getNativeUsage() {
Map<String, Integer> usage = new HashMap<>();
usage.put("promptTokens", getPromptTokens());
usage.put("completionTokens", getCompletionTokens());
usage.put("totalTokens", getTotalTokens());
return usage;
}
}
}

Some files were not shown because too many files have changed in this diff Show More