From cd886437eb1a6a5ec9bc7b7bf367a94a098c1dcc Mon Sep 17 00:00:00 2001 From: VictorZalevski Date: Wed, 6 Nov 2024 09:16:55 +0300 Subject: [PATCH] Improve usage field to include new properties OpenAI's API returns additional token usage metrics that provide deeper insight into API consumption. This adds support for: - acceptedPredictionTokens: Tokens from accepted model predictions - audioTokens: Tokens used for audio processing - rejectedPredictionTokens: Tokens from rejected model predictions These fields help track resource utilization and costs more accurately by breaking down token usage by type. Added @JsonIgnoreProperties to maintain compatibility with future OpenAI API additions. Fixes warning logging in RetryUtils.SHORT_RETRY_TEMPLATE to reduce noise in test output. --- .../ai/openai/api/OpenAiApi.java | 7 ++- .../ai/openai/metadata/OpenAiUsage.java | 20 +++++++ .../ai/openai/metadata/OpenAiUsageTests.java | 52 ++++++++++++++++++- .../text/VertexAiTextEmbeddingRetryTests.java | 2 +- .../springframework/ai/retry/RetryUtils.java | 5 +- 5 files changed, 80 insertions(+), 6 deletions(-) diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java index 20c0eeab6..cb0ca13ca 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java @@ -23,6 +23,7 @@ import java.util.function.Consumer; import java.util.function.Predicate; import com.fasterxml.jackson.annotation.JsonIgnore; +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; import com.fasterxml.jackson.annotation.JsonProperty; @@ -1253,8 +1254,12 @@ public class OpenAiApi { * @param reasoningTokens Number of tokens generated by the model for reasoning. */ @JsonInclude(Include.NON_NULL) + @JsonIgnoreProperties(ignoreUnknown = true) public record CompletionTokenDetails(// @formatter:off - @JsonProperty("reasoning_tokens") Integer reasoningTokens) { // @formatter:on + @JsonProperty("reasoning_tokens") Integer reasoningTokens, + @JsonProperty("accepted_prediction_tokens") Integer acceptedPredictionTokens, + @JsonProperty("audio_tokens") Integer audioTokens, + @JsonProperty("rejected_prediction_tokens") Integer rejectedPredictionTokens) { // @formatter:on } } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiUsage.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiUsage.java index 4e32bd153..e72e1dac1 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiUsage.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiUsage.java @@ -72,6 +72,26 @@ public class OpenAiUsage implements Usage { return reasoningTokens != null ? reasoningTokens.longValue() : 0; } + public Long getAcceptedPredictionTokens() { + OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails(); + Integer acceptedPredictionTokens = completionTokenDetails != null + ? completionTokenDetails.acceptedPredictionTokens() : null; + return acceptedPredictionTokens != null ? acceptedPredictionTokens.longValue() : 0; + } + + public Long getAudioTokens() { + OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails(); + Integer audioTokens = completionTokenDetails != null ? completionTokenDetails.audioTokens() : null; + return audioTokens != null ? audioTokens.longValue() : 0; + } + + public Long getRejectedPredictionTokens() { + OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails(); + Integer rejectedPredictionTokens = completionTokenDetails != null + ? completionTokenDetails.rejectedPredictionTokens() : null; + return rejectedPredictionTokens != null ? rejectedPredictionTokens.longValue() : 0; + } + @Override public Long getTotalTokens() { Integer totalTokens = getUsage().totalTokens(); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/metadata/OpenAiUsageTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/metadata/OpenAiUsageTests.java index 65a97c1c6..962af2e22 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/metadata/OpenAiUsageTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/metadata/OpenAiUsageTests.java @@ -83,7 +83,7 @@ class OpenAiUsageTests { @Test void whenReasoningTokensIsNull() { OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, - new OpenAiApi.Usage.CompletionTokenDetails(null)); + new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, null)); OpenAiUsage usage = OpenAiUsage.from(openAiUsage); assertThat(usage.getReasoningTokens()).isEqualTo(0); } @@ -91,11 +91,59 @@ class OpenAiUsageTests { @Test void whenCompletionTokenDetailsIsPresent() { OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, - new OpenAiApi.Usage.CompletionTokenDetails(50)); + new OpenAiApi.Usage.CompletionTokenDetails(50, null, null, null)); OpenAiUsage usage = OpenAiUsage.from(openAiUsage); assertThat(usage.getReasoningTokens()).isEqualTo(50); } + @Test + void whenAcceptedPredictionTokensIsNull() { + 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.getAcceptedPredictionTokens()).isEqualTo(0); + } + + @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.getAcceptedPredictionTokens()).isEqualTo(75); + } + + @Test + void whenAudioTokensIsNull() { + 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.getAudioTokens()).isEqualTo(0); + } + + @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.getAudioTokens()).isEqualTo(125); + } + + @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.getRejectedPredictionTokens()).isEqualTo(0); + } + + @Test + 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.getRejectedPredictionTokens()).isEqualTo(25); + } + @Test void whenCacheTokensIsNull() { OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, new OpenAiApi.Usage.PromptTokensDetails(null), diff --git a/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingRetryTests.java b/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingRetryTests.java index 9d2a2bd07..3430791d5 100644 --- a/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingRetryTests.java +++ b/models/spring-ai-vertex-ai-embedding/src/test/java/org/springframework/ai/vertexai/embedding/text/VertexAiTextEmbeddingRetryTests.java @@ -73,7 +73,7 @@ public class VertexAiTextEmbeddingRetryTests { @BeforeEach public void setUp() { - this.retryTemplate = RetryUtils.DEFAULT_RETRY_TEMPLATE; + this.retryTemplate = RetryUtils.SHORT_RETRY_TEMPLATE; this.retryListener = new TestRetryListener(); this.retryTemplate.registerListener(this.retryListener); diff --git a/spring-ai-retry/src/main/java/org/springframework/ai/retry/RetryUtils.java b/spring-ai-retry/src/main/java/org/springframework/ai/retry/RetryUtils.java index 1207bfb2c..960382d5c 100644 --- a/spring-ai-retry/src/main/java/org/springframework/ai/retry/RetryUtils.java +++ b/spring-ai-retry/src/main/java/org/springframework/ai/retry/RetryUtils.java @@ -84,7 +84,8 @@ public abstract class RetryUtils { .build(); /** - * Useful in testing scenarios where you don't want to wait long for retry. + * Useful in testing scenarios where you don't want to wait long for retry and now + * show stack trace */ public static final RetryTemplate SHORT_RETRY_TEMPLATE = RetryTemplate.builder() .maxAttempts(10) @@ -95,7 +96,7 @@ public abstract class RetryUtils { @Override public void onError(RetryContext context, RetryCallback callback, Throwable throwable) { - logger.warn("Retry error. Retry count:" + context.getRetryCount(), throwable); + logger.warn("Retry error. Retry count:" + context.getRetryCount()); } }) .build();