improved ability to use openAiApi to support DeepSeek

Signed-off-by: Ricken Bazolo <ricken.bazolo@gmail.com>
This commit is contained in:
Ricken Bazolo
2025-01-25 14:30:28 +01:00
committed by Ilayaperumal Gopinathan
parent 35101e75ae
commit 45421b1c93
2 changed files with 47 additions and 13 deletions

View File

@@ -1412,17 +1412,28 @@ public class OpenAiApi {
* completion).
* @param promptTokensDetails Breakdown of tokens used in the prompt.
* @param completionTokenDetails Breakdown of tokens used in a completion.
* @param promptCacheHitTokens Number of tokens in the prompt that were served from
* (util for
* <a href="https://api-docs.deepseek.com/api/create-chat-completion">DeepSeek</a>
* support).
* @param promptCacheMissTokens Number of tokens in the prompt that were not served
* (util for
* <a href="https://api-docs.deepseek.com/api/create-chat-completion">DeepSeek</a>
* support).
*/
@JsonInclude(Include.NON_NULL)
@JsonIgnoreProperties(ignoreUnknown = true)
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_details") CompletionTokenDetails completionTokenDetails,
@JsonProperty("prompt_cache_hit_tokens") Integer promptCacheHitTokens,
@JsonProperty("prompt_cache_miss_tokens") Integer promptCacheMissTokens) { // @formatter:on
public Usage(Integer completionTokens, Integer promptTokens, Integer totalTokens) {
this(completionTokens, promptTokens, totalTokens, null, null);
this(completionTokens, promptTokens, totalTokens, null, null, null, null);
}
/**

View File

@@ -19,7 +19,6 @@ 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;
@@ -81,7 +80,7 @@ class OpenAiUsageTests {
@Test
void whenPromptAndCompletionTokensDetailsIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, null);
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, null, null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getTotalTokens()).isEqualTo(300);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
@@ -91,7 +90,7 @@ class OpenAiUsageTests {
@Test
void whenCompletionTokenDetailsIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, null);
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null, null, null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
assertThat(usage.getTotalTokens()).isEqualTo(300);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
@@ -101,7 +100,7 @@ class OpenAiUsageTests {
@Test
void whenReasoningTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, null));
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, null), null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
@@ -110,7 +109,7 @@ class OpenAiUsageTests {
@Test
void whenCompletionTokenDetailsIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(50, null, null, null));
new OpenAiApi.Usage.CompletionTokenDetails(50, null, null, null), null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(50);
@@ -122,7 +121,7 @@ class OpenAiUsageTests {
@Test
void whenAcceptedPredictionTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, 75, null, null));
new OpenAiApi.Usage.CompletionTokenDetails(null, 75, null, null), null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
@@ -134,7 +133,7 @@ class OpenAiUsageTests {
@Test
void whenAudioTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, null, 125, null));
new OpenAiApi.Usage.CompletionTokenDetails(null, null, 125, null), null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
@@ -146,7 +145,7 @@ class OpenAiUsageTests {
@Test
void whenRejectedPredictionTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, null));
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, null), null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
@@ -160,7 +159,7 @@ class OpenAiUsageTests {
@Test
void whenRejectedPredictionTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300, null,
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, 25));
new OpenAiApi.Usage.CompletionTokenDetails(null, null, null, 25), null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.completionTokenDetails().reasoningTokens()).isEqualTo(null);
@@ -172,7 +171,7 @@ class OpenAiUsageTests {
@Test
void whenCacheTokensIsNull() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300,
new OpenAiApi.Usage.PromptTokensDetails(null, null), null);
new OpenAiApi.Usage.PromptTokensDetails(null, null), null, null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.promptTokensDetails().audioTokens()).isEqualTo(null);
@@ -182,11 +181,35 @@ class OpenAiUsageTests {
@Test
void whenCacheTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300,
new OpenAiApi.Usage.PromptTokensDetails(99, 15), null);
new OpenAiApi.Usage.PromptTokensDetails(99, 15), null, null, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.promptTokensDetails().audioTokens()).isEqualTo(99);
assertThat(nativeUsage.promptTokensDetails().cachedTokens()).isEqualTo(15);
}
@Test
void whenPromptCacheHitTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300,
new OpenAiApi.Usage.PromptTokensDetails(99, 15), null, 150, null);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.promptTokensDetails().audioTokens()).isEqualTo(99);
assertThat(nativeUsage.promptTokensDetails().cachedTokens()).isEqualTo(15);
assertThat(nativeUsage.promptCacheHitTokens()).isEqualTo(150);
assertThat(nativeUsage.promptCacheMissTokens()).isNull();
}
@Test
void whenPromptCacheMissTokensIsPresent() {
OpenAiApi.Usage openAiUsage = new OpenAiApi.Usage(100, 200, 300,
new OpenAiApi.Usage.PromptTokensDetails(99, 15), null, null, 80);
DefaultUsage usage = getDefaultUsage(openAiUsage);
OpenAiApi.Usage nativeUsage = (OpenAiApi.Usage) usage.getNativeUsage();
assertThat(nativeUsage.promptTokensDetails().audioTokens()).isEqualTo(99);
assertThat(nativeUsage.promptTokensDetails().cachedTokens()).isEqualTo(15);
assertThat(nativeUsage.promptCacheMissTokens()).isEqualTo(80);
assertThat(nativeUsage.promptCacheHitTokens()).isNull();
}
}