Add support for reasoning tokens in OpenAI usage data

This change introduces a new field for tracking reasoning tokens in the
OpenAI API response. It extends the Usage record to include
CompletionTokenDetails, allowing for more granular token usage
reporting. The OpenAiUsage class is updated to expose this new data,
and corresponding unit tests are added to verify the behavior.

This enhancement provides more detailed insights into token usage,
particularly for advanced AI models that separate reasoning from other
generation processes.
This commit is contained in:
dafriz
2024-09-22 21:53:20 +10:00
committed by Mark Pollack
parent 835450761e
commit 1673907db0
3 changed files with 47 additions and 1 deletions

View File

@@ -936,12 +936,28 @@ public class OpenAiApi {
* @param promptTokens Number of tokens in the prompt.
* @param totalTokens Total number of tokens used in the request (prompt +
* completion).
* @param completionTokenDetails Breakdown of tokens used in a completion
*/
@JsonInclude(Include.NON_NULL)
public record Usage(// @formatter:off
@JsonProperty("completion_tokens") Integer completionTokens,
@JsonProperty("prompt_tokens") Integer promptTokens,
@JsonProperty("total_tokens") Integer totalTokens) {// @formatter:on
@JsonProperty("total_tokens") Integer totalTokens,
@JsonProperty("completion_tokens_details") CompletionTokenDetails completionTokenDetails) {// @formatter:on
public Usage(Integer completionTokens, Integer promptTokens, Integer totalTokens) {
this(completionTokens, promptTokens, totalTokens, null);
}
/**
* Breakdown of tokens used in a completion
*
* @param reasoningTokens Number of tokens generated by the model for reasoning.
*/
@JsonInclude(Include.NON_NULL)
public record CompletionTokenDetails(// @formatter:off
@JsonProperty("reasoning_tokens") Integer reasoningTokens) {// @formatter:on
}
}

View File

@@ -58,6 +58,12 @@ public class OpenAiUsage implements Usage {
return generationTokens != null ? generationTokens.longValue() : 0;
}
public Long getReasoningTokens() {
OpenAiApi.Usage.CompletionTokenDetails completionTokenDetails = getUsage().completionTokenDetails();
Integer reasoningTokens = completionTokenDetails != null ? completionTokenDetails.reasoningTokens() : null;
return reasoningTokens != null ? reasoningTokens.longValue() : 0;
}
@Override
public Long getTotalTokens() {
Integer totalTokens = getUsage().totalTokens();

View File

@@ -54,4 +54,28 @@ class OpenAiUsageTests {
assertThat(usage.getTotalTokens()).isEqualTo(300);
}
@Test
void whenCompletionTokenDetailsIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null);
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getTotalTokens()).isEqualTo(300);
assertThat(usage.getReasoningTokens()).isEqualTo(0);
}
@Test
void whenReasoningTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300,
new OpenAiApi.Usage.CompletionTokenDetails(null));
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getReasoningTokens()).isEqualTo(0);
}
@Test
void whenCompletionTokenDetailsIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300,
new OpenAiApi.Usage.CompletionTokenDetails(50));
OpenAiUsage usage = OpenAiUsage.from(openAiUsage);
assertThat(usage.getReasoningTokens()).isEqualTo(50);
}
}