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:
committed by
Mark Pollack
parent
840304955e
commit
4b64aa0ca6
@@ -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) {
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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());
|
||||
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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));
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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()))
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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 + "]";
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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()) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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());
|
||||
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"));
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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);
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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()))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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
22
pom.xml
@@ -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>
|
||||
|
||||
|
||||
@@ -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 + '}';
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user