Restructure inclusion of data from Advisors in ChatResponse

* Make OpenAiChatResponseMetadata support JSON SerDeser
This commit is contained in:
Mark Pollack
2024-06-04 01:06:40 -04:00
parent ac91302eed
commit 59da8d37dc
7 changed files with 156 additions and 28 deletions

View File

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

View File

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

View File

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

View File

@@ -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()

View File

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

View File

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

View File

@@ -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 {
}