Restructure inclusion of data from Advisors in ChatResponse
* Make OpenAiChatResponseMetadata support JSON SerDeser
This commit is contained in:
@@ -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<String, Object> 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<String, Object> 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<String, Object> 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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -113,14 +113,14 @@ public class QuestionAnswerAdvisor implements RequestResponseAdvisor {
|
||||
|
||||
@Override
|
||||
public ChatResponse adviseResponse(ChatResponse response, Map<String, Object> context) {
|
||||
response.getMetadata().put(RETRIEVED_DOCUMENTS, context.get(RETRIEVED_DOCUMENTS));
|
||||
response.getAdvisorContext().put(RETRIEVED_DOCUMENTS, context.get(RETRIEVED_DOCUMENTS));
|
||||
return response;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatResponse> adviseResponse(Flux<ChatResponse> fluxResponse, Map<String, Object> 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;
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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<Generation> {
|
||||
*/
|
||||
private final List<Generation> generations;
|
||||
|
||||
private Map<String, Object> 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<Generation> generations) {
|
||||
this(generations, ChatResponseMetadata.NULL);
|
||||
this(generations, ChatResponseMetadata.NULL, new HashMap<>());
|
||||
}
|
||||
|
||||
public ChatResponse(List<Generation> generations, ChatResponseMetadata chatResponseMetadata) {
|
||||
this(generations, chatResponseMetadata, new HashMap<>());
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -50,9 +58,11 @@ public class ChatResponse implements ModelResponse<Generation> {
|
||||
* @param chatResponseMetadata {@link ChatResponseMetadata} containing information
|
||||
* about the use of the AI provider's API.
|
||||
*/
|
||||
public ChatResponse(List<Generation> generations, ChatResponseMetadata chatResponseMetadata) {
|
||||
this.chatResponseMetadata = chatResponseMetadata;
|
||||
public ChatResponse(List<Generation> generations, ChatResponseMetadata chatResponseMetadata,
|
||||
Map<String, Object> advisorContext) {
|
||||
this.generations = List.copyOf(generations);
|
||||
this.chatResponseMetadata = chatResponseMetadata;
|
||||
this.advisorContext = advisorContext;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -87,9 +97,14 @@ public class ChatResponse implements ModelResponse<Generation> {
|
||||
return this.chatResponseMetadata;
|
||||
}
|
||||
|
||||
public Map<String, Object> 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<Generation> {
|
||||
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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -27,6 +27,6 @@ import java.util.Map;
|
||||
* @author Mark Pollack
|
||||
* @since 0.8.0
|
||||
*/
|
||||
public interface ResponseMetadata extends Map<String, Object> {
|
||||
public interface ResponseMetadata {
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user