diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadata.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadata.java index f926496c1..461d090dd 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadata.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiChatResponseMetadata.java @@ -15,16 +15,20 @@ */ package org.springframework.ai.openai.metadata; +import com.fasterxml.jackson.annotation.JsonCreator; +import com.fasterxml.jackson.annotation.JsonIgnore; +import com.fasterxml.jackson.annotation.JsonProperty; import org.springframework.ai.chat.metadata.ChatResponseMetadata; import org.springframework.ai.chat.metadata.EmptyRateLimit; import org.springframework.ai.chat.metadata.EmptyUsage; +import org.springframework.ai.chat.metadata.PromptMetadata; import org.springframework.ai.chat.metadata.RateLimit; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.lang.Nullable; import org.springframework.util.Assert; -import java.util.HashMap; +import java.util.Objects; /** * {@link ChatResponseMetadata} implementation for {@literal OpenAI}. @@ -35,7 +39,7 @@ import java.util.HashMap; * @see Usage * @since 0.7.0 */ -public class OpenAiChatResponseMetadata extends HashMap implements ChatResponseMetadata { +public class OpenAiChatResponseMetadata implements ChatResponseMetadata { protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, usage: %3$s, rateLimit: %4$s }"; @@ -46,18 +50,26 @@ public class OpenAiChatResponseMetadata extends HashMap implemen return chatResponseMetadata; } - private final String id; + @JsonProperty("id") + private String id; @Nullable + @JsonProperty("rateLimit") private RateLimit rateLimit; - private final Usage usage; + @JsonProperty("usage") + private Usage usage; - protected OpenAiChatResponseMetadata(String id, OpenAiUsage usage) { + @JsonIgnore + private PromptMetadata promptMetadata; + + public OpenAiChatResponseMetadata(String id, OpenAiUsage usage) { this(id, usage, null); } - protected OpenAiChatResponseMetadata(String id, OpenAiUsage usage, @Nullable OpenAiRateLimit rateLimit) { + @JsonCreator + public OpenAiChatResponseMetadata(@JsonProperty("id") String id, @JsonProperty("usage") OpenAiUsage usage, + @Nullable @JsonProperty("rateLimit") OpenAiRateLimit rateLimit) { this.id = id; this.usage = usage; this.rateLimit = rateLimit; @@ -85,9 +97,29 @@ public class OpenAiChatResponseMetadata extends HashMap implemen return this; } + @Override + public PromptMetadata getPromptMetadata() { + return PromptMetadata.empty(); + } + @Override public String toString() { return AI_METADATA_STRING.formatted(getClass().getName(), getId(), getUsage(), getRateLimit()); } + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (!(o instanceof OpenAiChatResponseMetadata that)) + return false; + return Objects.equals(id, that.id) && Objects.equals(rateLimit, that.rateLimit) + && Objects.equals(usage, that.usage) && Objects.equals(promptMetadata, that.promptMetadata); + } + + @Override + public int hashCode() { + return Objects.hash(id, rateLimit, usage, promptMetadata); + } + } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiRateLimit.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiRateLimit.java index 7f5f214da..5c18ae07b 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiRateLimit.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/metadata/OpenAiRateLimit.java @@ -16,7 +16,10 @@ package org.springframework.ai.openai.metadata; import java.time.Duration; +import java.util.Objects; +import com.fasterxml.jackson.annotation.JsonCreator; +import com.fasterxml.jackson.annotation.JsonProperty; import org.springframework.ai.chat.metadata.RateLimit; /** @@ -44,8 +47,11 @@ public class OpenAiRateLimit implements RateLimit { private final Duration tokensReset; - public OpenAiRateLimit(Long requestsLimit, Long requestsRemaining, Duration requestsReset, Long tokensLimit, - Long tokensRemaining, Duration tokensReset) { + @JsonCreator + public OpenAiRateLimit(@JsonProperty("requestsLimit") Long requestsLimit, + @JsonProperty("requestsRemaining") Long requestsRemaining, + @JsonProperty("requestsReset") Duration requestsReset, @JsonProperty("tokensLimit") Long tokensLimit, + @JsonProperty("tokensRemaining") Long tokensRemaining, @JsonProperty("tokensReset") Duration tokensReset) { this.requestsLimit = requestsLimit; this.requestsRemaining = requestsRemaining; @@ -91,4 +97,22 @@ public class OpenAiRateLimit implements RateLimit { getRequestsReset(), getTokensLimit(), getTokensRemaining(), getTokensReset()); } + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (!(o instanceof OpenAiRateLimit that)) + return false; + return Objects.equals(requestsLimit, that.requestsLimit) + && Objects.equals(requestsRemaining, that.requestsRemaining) + && Objects.equals(tokensLimit, that.tokensLimit) + && Objects.equals(tokensRemaining, that.tokensRemaining) + && Objects.equals(requestsReset, that.requestsReset) && Objects.equals(tokensReset, that.tokensReset); + } + + @Override + public int hashCode() { + return Objects.hash(requestsLimit, requestsRemaining, tokensLimit, tokensRemaining, requestsReset, tokensReset); + } + } 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 5f1367736..6ce8281c2 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 @@ -15,10 +15,14 @@ */ package org.springframework.ai.openai.metadata; +import com.fasterxml.jackson.annotation.JsonCreator; +import com.fasterxml.jackson.annotation.JsonProperty; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.util.Assert; +import java.util.Objects; + /** * {@link Usage} implementation for {@literal OpenAI}. * @@ -34,35 +38,62 @@ public class OpenAiUsage implements Usage { return new OpenAiUsage(usage); } - private final OpenAiApi.Usage usage; + private Long promptTokens; - protected OpenAiUsage(OpenAiApi.Usage usage) { - Assert.notNull(usage, "OpenAI Usage must not be null"); - this.usage = usage; + private Long generationTokens; + + private Long totalTokens; + + public OpenAiUsage(OpenAiApi.Usage usage) { + Assert.notNull(usage, "OpenAiApi.Usage must not be null"); + this.promptTokens = usage.promptTokens().longValue(); + this.generationTokens = usage.completionTokens().longValue(); + this.totalTokens = usage.totalTokens().longValue(); } - protected OpenAiApi.Usage getUsage() { - return this.usage; + @JsonCreator + public OpenAiUsage(@JsonProperty("promptTokens") Long promptTokens, + @JsonProperty("generationTokens") Long generationTokens, @JsonProperty("totalTokens") Long totalTokens) { + this.promptTokens = promptTokens; + this.generationTokens = generationTokens; + this.totalTokens = totalTokens; } @Override public Long getPromptTokens() { - return getUsage().promptTokens().longValue(); + return this.promptTokens; } @Override public Long getGenerationTokens() { - return getUsage().completionTokens().longValue(); + return this.generationTokens; } @Override public Long getTotalTokens() { - return getUsage().totalTokens().longValue(); + return this.totalTokens; } @Override public String toString() { - return getUsage().toString(); + return "OpenAiUsage{" + "promptTokens=" + promptTokens + ", generationTokens=" + generationTokens + + ", totalTokens=" + totalTokens + '}'; + } + + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (!(o instanceof OpenAiUsage that)) + return false; + return Objects.equals(promptTokens, that.promptTokens) + && Objects.equals(generationTokens, that.generationTokens) + && Objects.equals(totalTokens, that.totalTokens); + } + + @Override + public int hashCode() { + return Objects.hash(promptTokens, generationTokens, totalTokens); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/client/OpenAiChatClientIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/client/OpenAiChatClientIT.java index 830cd0f7d..936fa0d77 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/client/OpenAiChatClientIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/client/OpenAiChatClientIT.java @@ -17,17 +17,27 @@ package org.springframework.ai.openai.chat.client; import java.io.IOException; import java.net.URL; +import java.time.Duration; import java.util.Arrays; import java.util.List; import java.util.Map; import java.util.stream.Collectors; +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.SerializationFeature; +import com.fasterxml.jackson.datatype.jsr310.JavaTimeModule; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.ValueSource; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import org.springframework.ai.chat.metadata.ChatResponseMetadata; +import org.springframework.ai.chat.metadata.PromptMetadata; +import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadata; +import org.springframework.ai.openai.metadata.OpenAiRateLimit; +import org.springframework.ai.openai.metadata.OpenAiUsage; import reactor.core.publisher.Flux; import org.springframework.ai.chat.client.ChatClient; @@ -62,7 +72,23 @@ class OpenAiChatClientIT extends AbstractIT { } @Test - void call() { + void serDeserChatResponseMetadata() throws JsonProcessingException { + OpenAiUsage openAiUsage = new OpenAiUsage(new OpenAiApi.Usage(1, 2, 3)); + OpenAiRateLimit openAiRateLimit = new OpenAiRateLimit(1L, 2L, Duration.ZERO, 4L, 5L, Duration.ZERO); + OpenAiChatResponseMetadata chatResponseMetadata = new OpenAiChatResponseMetadata("myid", openAiUsage, + openAiRateLimit); + ObjectMapper objectMapper = new ObjectMapper(); + objectMapper.enable(SerializationFeature.INDENT_OUTPUT); + objectMapper.registerModule(new JavaTimeModule()); + String json = objectMapper.writeValueAsString(chatResponseMetadata); + System.out.println("ChatResponseMetadata Ser: " + json); + + OpenAiChatResponseMetadata deserialized = objectMapper.readValue(json, OpenAiChatResponseMetadata.class); + assertThat(chatResponseMetadata).usingRecursiveComparison().isEqualTo(deserialized); + } + + @Test + void call() throws JsonProcessingException { // @formatter:off ChatResponse response = ChatClient.builder(chatModel).build().prompt() diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/client/advisor/QuestionAnswerAdvisor.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/client/advisor/QuestionAnswerAdvisor.java index 25df8dfca..ce7d32a24 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/client/advisor/QuestionAnswerAdvisor.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/client/advisor/QuestionAnswerAdvisor.java @@ -113,14 +113,14 @@ public class QuestionAnswerAdvisor implements RequestResponseAdvisor { @Override public ChatResponse adviseResponse(ChatResponse response, Map context) { - response.getMetadata().put(RETRIEVED_DOCUMENTS, context.get(RETRIEVED_DOCUMENTS)); + response.getAdvisorContext().put(RETRIEVED_DOCUMENTS, context.get(RETRIEVED_DOCUMENTS)); return response; } @Override public Flux adviseResponse(Flux fluxResponse, Map context) { return fluxResponse.map(cr -> { - cr.getMetadata().put(RETRIEVED_DOCUMENTS, context.get(RETRIEVED_DOCUMENTS)); + cr.getAdvisorContext().put(RETRIEVED_DOCUMENTS, context.get(RETRIEVED_DOCUMENTS)); return cr; }); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ChatResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ChatResponse.java index 8c31776c0..1c87fc132 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ChatResponse.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ChatResponse.java @@ -15,7 +15,9 @@ */ package org.springframework.ai.chat.model; +import java.util.HashMap; import java.util.List; +import java.util.Map; import java.util.Objects; import org.springframework.ai.model.ModelResponse; @@ -34,13 +36,19 @@ public class ChatResponse implements ModelResponse { */ private final List generations; + private Map advisorContext; + /** * Construct a new {@link ChatResponse} instance without metadata. * @param generations the {@link List} of {@link Generation} returned by the AI * provider. */ public ChatResponse(List generations) { - this(generations, ChatResponseMetadata.NULL); + this(generations, ChatResponseMetadata.NULL, new HashMap<>()); + } + + public ChatResponse(List generations, ChatResponseMetadata chatResponseMetadata) { + this(generations, chatResponseMetadata, new HashMap<>()); } /** @@ -50,9 +58,11 @@ public class ChatResponse implements ModelResponse { * @param chatResponseMetadata {@link ChatResponseMetadata} containing information * about the use of the AI provider's API. */ - public ChatResponse(List generations, ChatResponseMetadata chatResponseMetadata) { - this.chatResponseMetadata = chatResponseMetadata; + public ChatResponse(List generations, ChatResponseMetadata chatResponseMetadata, + Map advisorContext) { this.generations = List.copyOf(generations); + this.chatResponseMetadata = chatResponseMetadata; + this.advisorContext = advisorContext; } /** @@ -87,9 +97,14 @@ public class ChatResponse implements ModelResponse { return this.chatResponseMetadata; } + public Map getAdvisorContext() { + return this.advisorContext; + } + @Override public String toString() { - return "ChatResponse [metadata=" + chatResponseMetadata + ", generations=" + generations + "]"; + return "ChatResponse{" + "chatResponseMetadata=" + chatResponseMetadata + ", generations=" + generations + + ", advisorContext=" + advisorContext + '}'; } @Override @@ -99,12 +114,12 @@ public class ChatResponse implements ModelResponse { if (!(o instanceof ChatResponse that)) return false; return Objects.equals(chatResponseMetadata, that.chatResponseMetadata) - && Objects.equals(generations, that.generations); + && Objects.equals(generations, that.generations) && Objects.equals(advisorContext, that.advisorContext); } @Override public int hashCode() { - return Objects.hash(chatResponseMetadata, generations); + return Objects.hash(chatResponseMetadata, generations, advisorContext); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/ResponseMetadata.java b/spring-ai-core/src/main/java/org/springframework/ai/model/ResponseMetadata.java index b1516ec3f..151cc5e91 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/ResponseMetadata.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/ResponseMetadata.java @@ -27,6 +27,6 @@ import java.util.Map; * @author Mark Pollack * @since 0.8.0 */ -public interface ResponseMetadata extends Map { +public interface ResponseMetadata { }