SerDeser of ChatResponse

This commit is contained in:
Mark Pollack
2024-06-06 13:53:11 -04:00
parent 59da8d37dc
commit cf0946b891
17 changed files with 234 additions and 44 deletions

View File

@@ -18,6 +18,7 @@ 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 com.fasterxml.jackson.annotation.JsonTypeName;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.EmptyRateLimit;
import org.springframework.ai.chat.metadata.EmptyUsage;
@@ -39,6 +40,7 @@ import java.util.Objects;
* @see Usage
* @since 0.7.0
*/
@JsonTypeName("openai")
public class OpenAiChatResponseMetadata implements ChatResponseMetadata {
protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, usage: %3$s, rateLimit: %4$s }";

View File

@@ -0,0 +1,68 @@
/*
* Copyright 2024 - 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.openai.chat;
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.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadata;
import org.springframework.ai.openai.metadata.OpenAiRateLimit;
import org.springframework.ai.openai.metadata.OpenAiUsage;
import java.time.Duration;
import static org.assertj.core.api.Assertions.assertThat;
public class OpenAiChatResponseTests {
@Test
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 promptSerialization() throws JsonProcessingException {
Prompt prompt = new Prompt(new UserMessage("hello world"));
ObjectMapper objectMapper = new ObjectMapper();
objectMapper.enable(SerializationFeature.INDENT_OUTPUT);
String json = objectMapper.writeValueAsString(prompt);
System.out.println("Prompt Ser: " + json);
Prompt deserializedPrompt = objectMapper.readValue(json, Prompt.class);
assertThat(prompt).usingRecursiveComparison().isEqualTo(deserializedPrompt);
}
}

View File

@@ -71,22 +71,6 @@ class OpenAiChatClientIT extends AbstractIT {
record ActorsFilms(String actor, List<String> movies) {
}
@Test
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 {

View File

@@ -26,6 +26,8 @@ import java.util.List;
import java.util.Map;
import java.util.Objects;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.core.io.Resource;
import org.springframework.util.Assert;
import org.springframework.util.StreamUtils;
@@ -69,8 +71,10 @@ public abstract class AbstractMessage implements Message {
this(messageType, textContent, media, Map.of(MESSAGE_TYPE, messageType));
}
protected AbstractMessage(MessageType messageType, String textContent, Collection<Media> media,
Map<String, Object> metadata) {
@JsonCreator
protected AbstractMessage(@JsonProperty("messageType") MessageType messageType,
@JsonProperty("content") String textContent, @JsonProperty("media") Collection<Media> media,
@JsonProperty("metadata") Map<String, Object> metadata) {
Assert.notNull(messageType, "Message type must not be null");
Assert.notNull(textContent, "Content must not be null");

View File

@@ -15,6 +15,9 @@
*/
package org.springframework.ai.chat.messages;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonProperty;
import java.util.Map;
/**
@@ -25,11 +28,13 @@ import java.util.Map;
*/
public class AssistantMessage extends AbstractMessage {
public AssistantMessage(String content) {
super(MessageType.ASSISTANT, content);
public AssistantMessage(String textContent) {
super(MessageType.ASSISTANT, textContent);
}
public AssistantMessage(String content, Map<String, Object> properties) {
@JsonCreator
public AssistantMessage(@JsonProperty("content") String content,
@JsonProperty("metadata") Map<String, Object> properties) {
super(MessageType.ASSISTANT, content, properties);
}

View File

@@ -15,6 +15,9 @@
*/
package org.springframework.ai.chat.messages;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonProperty;
import java.util.Map;
/**
@@ -27,7 +30,9 @@ public class FunctionMessage extends AbstractMessage {
super(MessageType.FUNCTION, content);
}
public FunctionMessage(String content, Map<String, Object> properties) {
@JsonCreator
public FunctionMessage(@JsonProperty("content") String content,
@JsonProperty("metadata") Map<String, Object> properties) {
super(MessageType.FUNCTION, content, properties);
}

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.chat.messages;
import com.fasterxml.jackson.annotation.JsonSubTypes;
import com.fasterxml.jackson.annotation.JsonTypeInfo;
import org.springframework.ai.model.Content;
/**
@@ -25,6 +27,11 @@ import org.springframework.ai.model.Content;
* @see Media
* @see MessageType
*/
@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, include = JsonTypeInfo.As.PROPERTY, property = "messageType")
@JsonSubTypes({ @JsonSubTypes.Type(value = UserMessage.class, name = "USER"),
@JsonSubTypes.Type(value = SystemMessage.class, name = "SYSTEM"),
@JsonSubTypes.Type(value = AssistantMessage.class, name = "ASSISTANT"),
@JsonSubTypes.Type(value = FunctionMessage.class, name = "FUNCTION") })
public interface Message extends Content {
MessageType getMessageType();

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.chat.messages;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.core.io.Resource;
/**
@@ -26,7 +28,8 @@ import org.springframework.core.io.Resource;
*/
public class SystemMessage extends AbstractMessage {
public SystemMessage(String content) {
@JsonCreator
public SystemMessage(@JsonProperty("content") String content) {
super(MessageType.SYSTEM, content);
}

View File

@@ -20,6 +20,8 @@ import java.util.Collection;
import java.util.List;
import java.util.Map;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.core.io.Resource;
/**
@@ -45,8 +47,10 @@ public class UserMessage extends AbstractMessage {
this(textContent, Arrays.asList(media));
}
public UserMessage(String textContent, Collection<Media> mediaList, Map<String, Object> metadata) {
super(MessageType.USER, textContent, mediaList, metadata);
@JsonCreator
public UserMessage(@JsonProperty("content") String content, @JsonProperty("media") Collection<Media> mediaList,
@JsonProperty("metadata") Map<String, Object> metadata) {
super(MessageType.USER, content, mediaList, metadata);
}
@Override

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.chat.metadata;
import com.fasterxml.jackson.annotation.JsonSubTypes;
import com.fasterxml.jackson.annotation.JsonTypeInfo;
import org.springframework.ai.model.ResultMetadata;
import org.springframework.lang.Nullable;
@@ -25,6 +27,8 @@ import org.springframework.lang.Nullable;
* @author John Blum
* @since 0.7.0
*/
@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, include = JsonTypeInfo.As.PROPERTY, property = "type")
@JsonSubTypes({ @JsonSubTypes.Type(value = DefaultChatGenerationMetadata.class, name = "default") })
public interface ChatGenerationMetadata extends ResultMetadata {
ChatGenerationMetadata NULL = ChatGenerationMetadata.from(null, null);
@@ -39,19 +43,7 @@ public interface ChatGenerationMetadata extends ResultMetadata {
* reason} and content filter metadata.
*/
static ChatGenerationMetadata from(String finishReason, Object contentFilterMetadata) {
return new ChatGenerationMetadata() {
@Override
@SuppressWarnings("unchecked")
public <T> T getContentFilterMetadata() {
return (T) contentFilterMetadata;
}
@Override
public String getFinishReason() {
return finishReason;
}
};
return new DefaultChatGenerationMetadata(finishReason, contentFilterMetadata);
}
/**

View File

@@ -15,6 +15,7 @@
*/
package org.springframework.ai.chat.metadata;
import com.fasterxml.jackson.annotation.JsonTypeInfo;
import org.springframework.ai.model.ResponseMetadata;
import java.util.HashMap;
@@ -26,6 +27,8 @@ import java.util.HashMap;
* @author John Blum
* @since 0.7.0
*/
@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, include = JsonTypeInfo.As.PROPERTY, property = "type")
public interface ChatResponseMetadata extends ResponseMetadata {
static class DefaultChatResponseMetadata extends HashMap<String, Object> implements ChatResponseMetadata {

View File

@@ -0,0 +1,55 @@
package org.springframework.ai.chat.metadata;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonProperty;
import java.util.Objects;
public class DefaultChatGenerationMetadata implements ChatGenerationMetadata {
private final String finishReason;
private final Object contentFilterMetadata;
@JsonCreator
public DefaultChatGenerationMetadata(@JsonProperty("finishReason") String finishReason,
@JsonProperty("contentFilterMetadata") Object contentFilterMetadata) {
this.finishReason = finishReason;
this.contentFilterMetadata = contentFilterMetadata;
}
@Override
@JsonProperty("contentFilterMetadata")
public <T> T getContentFilterMetadata() {
return (T) contentFilterMetadata;
}
@Override
@JsonProperty("finishReason")
public String getFinishReason() {
return finishReason;
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (!(o instanceof DefaultChatGenerationMetadata))
return false;
DefaultChatGenerationMetadata that = (DefaultChatGenerationMetadata) o;
return Objects.equals(finishReason, that.finishReason)
&& Objects.equals(contentFilterMetadata, that.contentFilterMetadata);
}
@Override
public int hashCode() {
return Objects.hash(finishReason, contentFilterMetadata);
}
@Override
public String toString() {
return "DefaultChatGenerationMetadata{" + "finishReason='" + finishReason + '\'' + ", contentFilterMetadata="
+ contentFilterMetadata + '}';
}
}

View File

@@ -15,11 +15,15 @@
*/
package org.springframework.ai.chat.model;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonIgnore;
import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.ai.model.ModelResponse;
import org.springframework.util.CollectionUtils;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
@@ -34,9 +38,9 @@ public class ChatResponse implements ModelResponse<Generation> {
/**
* List of generated messages returned by the AI provider.
*/
private final List<Generation> generations;
private List<Generation> generations = new ArrayList<>();
private Map<String, Object> advisorContext;
private Map<String, Object> advisorContext = new HashMap<>();
/**
* Construct a new {@link ChatResponse} instance without metadata.
@@ -58,9 +62,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,
Map<String, Object> advisorContext) {
this.generations = List.copyOf(generations);
@JsonCreator
public ChatResponse(@JsonProperty("results") List<Generation> generations,
@JsonProperty("chatResponseMetadata") ChatResponseMetadata chatResponseMetadata,
@JsonProperty("advisorContext") Map<String, Object> advisorContext) {
this.generations = generations;
this.chatResponseMetadata = chatResponseMetadata;
this.advisorContext = advisorContext;
}
@@ -81,6 +87,7 @@ public class ChatResponse implements ModelResponse<Generation> {
/**
* @return Returns the first {@link Generation} in the generations list.
*/
@JsonIgnore
public Generation getResult() {
if (CollectionUtils.isEmpty(this.generations)) {
return null;

View File

@@ -18,6 +18,8 @@ package org.springframework.ai.chat.model;
import java.util.Map;
import java.util.Objects;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.ai.model.ModelResult;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.messages.AssistantMessage;
@@ -32,6 +34,13 @@ public class Generation implements ModelResult<AssistantMessage> {
private ChatGenerationMetadata chatGenerationMetadata;
@JsonCreator
public Generation(@JsonProperty("assistantMessage") AssistantMessage assistantMessage,
@JsonProperty("chatGenerationMetadata") ChatGenerationMetadata chatGenerationMetadata) {
this.assistantMessage = assistantMessage;
this.chatGenerationMetadata = chatGenerationMetadata;
}
public Generation(String text) {
this.assistantMessage = new AssistantMessage(text);
}
@@ -41,11 +50,13 @@ public class Generation implements ModelResult<AssistantMessage> {
}
@Override
@JsonProperty("assistantMessage")
public AssistantMessage getOutput() {
return this.assistantMessage;
}
@Override
@JsonProperty("chatGenerationMetadata")
public ChatGenerationMetadata getMetadata() {
ChatGenerationMetadata chatGenerationMetadata = this.chatGenerationMetadata;
return chatGenerationMetadata != null ? chatGenerationMetadata : ChatGenerationMetadata.NULL;

View File

@@ -20,6 +20,9 @@ import java.util.Collections;
import java.util.List;
import java.util.Objects;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonIgnore;
import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.FunctionMessage;
import org.springframework.ai.chat.messages.Message;
@@ -54,11 +57,14 @@ public class Prompt implements ModelRequest<List<Message>> {
this(Collections.singletonList(message), modelOptions);
}
public Prompt(List<Message> messages, ChatOptions modelOptions) {
@JsonCreator
public Prompt(@JsonProperty("instructions") List<Message> messages,
@JsonProperty("options") ChatOptions modelOptions) {
this.messages = messages;
this.modelOptions = modelOptions;
}
@JsonIgnore
public String getContents() {
StringBuilder sb = new StringBuilder();
for (Message message : getInstructions()) {

View File

@@ -1,5 +1,6 @@
package org.springframework.ai.model;
import com.fasterxml.jackson.annotation.JsonTypeInfo;
import org.springframework.ai.chat.messages.Media;
import java.util.Collection;
@@ -15,6 +16,7 @@ import java.util.Map;
* @author Christian Tzolov
* @since 1.0.0
*/
@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, include = JsonTypeInfo.As.PROPERTY, property = "type")
public interface Content {
/**

View File

@@ -0,0 +1,32 @@
package org.springframework.ai.chat.client;
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 static org.assertj.core.api.Assertions.assertThat;
public class AdvisedRequestTests {
@Test
void serDeserAdvisedRequest() throws JsonProcessingException {
AdvisedRequest.Builder builder = AdvisedRequest.builder();
AdvisedRequest advisedRequest = builder.withSystemText("This is system text")
.withUserText("This is user text")
.build();
ObjectMapper objectMapper = new ObjectMapper();
objectMapper.enable(SerializationFeature.INDENT_OUTPUT);
objectMapper.registerModule(new JavaTimeModule());
String json = objectMapper.writeValueAsString(advisedRequest);
System.out.println("AdvisedRequest Ser: " + json);
AdvisedRequest deserialized = objectMapper.readValue(json, AdvisedRequest.class);
assertThat(advisedRequest).usingRecursiveComparison().isEqualTo(deserialized);
}
}