From c038526dd11f3005a67f4e26b6222c3d64f02d87 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Mon, 18 Nov 2024 23:10:23 +0100 Subject: [PATCH] refactor(openai): consolidate token usage details and add audio tokens support The commit restructures OpenAI token usage tracking by: - Adding audio_tokens support in PromptTokensDetails - Deprecating individual token getter methods in favor of consolidated records - Introducing new PromptTokensDetails and CompletionTokenDetails records - Updating tests to reflect the new structure Resolves #1369 , #1720 --- .../ai/openai/api/OpenAiApi.java | 36 +++--- .../ai/openai/metadata/OpenAiUsage.java | 119 +++++++++++++----- .../ai/openai/metadata/OpenAiUsageTests.java | 71 +++++------ 3 files changed, 140 insertions(+), 86 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 119a0cc11..7b2dba5c2 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 @@ -1145,11 +1145,11 @@ public class OpenAiApi { */ @JsonInclude(Include.NON_NULL) public record Usage(// @formatter:off - @JsonProperty("completion_tokens") Integer completionTokens, - @JsonProperty("prompt_tokens") Integer promptTokens, - @JsonProperty("total_tokens") Integer totalTokens, - @JsonProperty("prompt_tokens_details") PromptTokensDetails promptTokensDetails, - @JsonProperty("completion_tokens_details") CompletionTokenDetails completionTokenDetails) { // @formatter:on + @JsonProperty("completion_tokens") Integer completionTokens, + @JsonProperty("prompt_tokens") Integer promptTokens, + @JsonProperty("total_tokens") Integer totalTokens, + @JsonProperty("prompt_tokens_details") PromptTokensDetails promptTokensDetails, + @JsonProperty("completion_tokens_details") CompletionTokenDetails completionTokenDetails) { // @formatter:on public Usage(Integer completionTokens, Integer promptTokens, Integer totalTokens) { this(completionTokens, promptTokens, totalTokens, null, null); @@ -1158,11 +1158,13 @@ public class OpenAiApi { /** * Breakdown of tokens used in the prompt * + * @param audioTokens Audio input tokens present in the prompt. * @param cachedTokens Cached tokens present in the prompt. */ @JsonInclude(Include.NON_NULL) public record PromptTokensDetails(// @formatter:off - @JsonProperty("cached_tokens") Integer cachedTokens) { // @formatter:on + @JsonProperty("audio_tokens") Integer audioTokens, + @JsonProperty("cached_tokens") Integer cachedTokens) { // @formatter:on } /** @@ -1178,10 +1180,10 @@ public class OpenAiApi { @JsonInclude(Include.NON_NULL) @JsonIgnoreProperties(ignoreUnknown = true) public record CompletionTokenDetails(// @formatter:off - @JsonProperty("reasoning_tokens") Integer reasoningTokens, - @JsonProperty("accepted_prediction_tokens") Integer acceptedPredictionTokens, - @JsonProperty("audio_tokens") Integer audioTokens, - @JsonProperty("rejected_prediction_tokens") Integer rejectedPredictionTokens) { // @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 } } @@ -1205,13 +1207,13 @@ public class OpenAiApi { */ @JsonInclude(Include.NON_NULL) public record ChatCompletionChunk(// @formatter:off - @JsonProperty("id") String id, - @JsonProperty("choices") List choices, - @JsonProperty("created") Long created, - @JsonProperty("model") String model, - @JsonProperty("system_fingerprint") String systemFingerprint, - @JsonProperty("object") String object, - @JsonProperty("usage") Usage usage) { // @formatter:on + @JsonProperty("id") String id, + @JsonProperty("choices") List choices, + @JsonProperty("created") Long created, + @JsonProperty("model") String model, + @JsonProperty("system_fingerprint") String systemFingerprint, + @JsonProperty("object") String object, + @JsonProperty("usage") Usage usage) { // @formatter:on /** * Chat completion choice. 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 e72e1dac1..14f429ff5 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 @@ -26,6 +26,7 @@ import org.springframework.util.Assert; * @author John Blum * @author Thomas Vitale * @author David Frizelle + * @author Christian Tzolov * @since 0.7.0 * @see Completion @@ -60,38 +61,6 @@ public class OpenAiUsage implements Usage { return generationTokens != null ? generationTokens.longValue() : 0; } - public Long getCachedTokens() { - OpenAiApi.Usage.PromptTokensDetails promptTokenDetails = getUsage().promptTokensDetails(); - Integer cachedTokens = promptTokenDetails != null ? promptTokenDetails.cachedTokens() : null; - return cachedTokens != null ? cachedTokens.longValue() : 0; - } - - public Long getReasoningTokens() { - OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails(); - Integer reasoningTokens = completionTokenDetails != null ? completionTokenDetails.reasoningTokens() : null; - 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(); @@ -103,9 +72,95 @@ public class OpenAiUsage implements Usage { } } + /** + * @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())); + } + + public record PromptTokensDetails(// @formatter:off + Integer audioTokens, + Integer cachedTokens) { + } + + public record CompletionTokenDetails( + Integer reasoningTokens, + Integer acceptedPredictionTokens, + Integer audioTokens, + Integer rejectedPredictionTokens) { // @formatter:on + } + @Override public String toString() { return getUsage().toString(); } + private int valueOrZero(Integer value) { + return value != null ? value : 0; + } + } 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 962af2e22..806af6c61 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 @@ -26,6 +26,7 @@ import static org.assertj.core.api.Assertions.assertThat; * Unit tests for {@link OpenAiUsage}. * * @author Thomas Vitale + * @author Christian Tzolov */ class OpenAiUsageTests { @@ -76,16 +77,10 @@ class OpenAiUsageTests { OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, null); OpenAiUsage usage = OpenAiUsage.from(openAiUsage); assertThat(usage.getTotalTokens()).isEqualTo(300); - assertThat(usage.getCachedTokens()).isEqualTo(0); - assertThat(usage.getReasoningTokens()).isEqualTo(0); - } - - @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); + assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0); } @Test @@ -93,15 +88,10 @@ class OpenAiUsageTests { 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.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); + assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(50); + assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0); } @Test @@ -109,15 +99,10 @@ class OpenAiUsageTests { 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); + assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(75); + assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0); } @Test @@ -125,7 +110,10 @@ class OpenAiUsageTests { 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); + assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(125); + assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0); } @Test @@ -133,7 +121,11 @@ class OpenAiUsageTests { 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); + assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(0); + } @Test @@ -141,23 +133,28 @@ class OpenAiUsageTests { 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); + assertThat(usage.getCompletionTokenDetails().reasoningTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().acceptedPredictionTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().audioTokens()).isEqualTo(0); + assertThat(usage.getCompletionTokenDetails().rejectedPredictionTokens()).isEqualTo(25); } @Test void whenCacheTokensIsNull() { - OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, new OpenAiApi.Usage.PromptTokensDetails(null), - null); + OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, + new OpenAiApi.Usage.PromptTokensDetails(null, null), null); OpenAiUsage usage = OpenAiUsage.from(openAiUsage); - assertThat(usage.getCachedTokens()).isEqualTo(0); + assertThat(usage.getPromptTokensDetails().audioTokens()).isEqualTo(0); + assertThat(usage.getPromptTokensDetails().cachedTokens()).isEqualTo(0); } @Test void whenCacheTokensIsPresent() { - OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, new OpenAiApi.Usage.PromptTokensDetails(15), - null); + OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, + new OpenAiApi.Usage.PromptTokensDetails(99, 15), null); OpenAiUsage usage = OpenAiUsage.from(openAiUsage); - assertThat(usage.getCachedTokens()).isEqualTo(15); + assertThat(usage.getPromptTokensDetails().audioTokens()).isEqualTo(99); + assertThat(usage.getPromptTokensDetails().cachedTokens()).isEqualTo(15); } }