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
This commit is contained in:
Christian Tzolov
2024-11-18 23:10:23 +01:00
parent fb2e7528d0
commit c038526dd1
3 changed files with 140 additions and 86 deletions

View File

@@ -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<ChunkChoice> 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<ChunkChoice> 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.

View File

@@ -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 <a href=
* "https://platform.openai.com/docs/api-reference/completions/object">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;
}
}

View File

@@ -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);
}
}