more fixes for openai ser-deser

This commit is contained in:
Mark Pollack
2024-06-08 19:34:33 -04:00
parent 2c4554eb3f
commit 5594cc4bef
3 changed files with 30 additions and 11 deletions

View File

@@ -15,17 +15,11 @@
*/
package org.springframework.ai.openai;
import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
import com.fasterxml.jackson.annotation.JsonIgnore;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonInclude.Include;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.annotation.JsonTypeName;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallingOptions;
@@ -35,11 +29,18 @@ import org.springframework.ai.openai.api.OpenAiApi.FunctionTool;
import org.springframework.boot.context.properties.NestedConfigurationProperty;
import org.springframework.util.Assert;
import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
/**
* @author Christian Tzolov
* @since 0.8.0
*/
@JsonInclude(Include.NON_NULL)
@JsonTypeName("OpenAiChatOptions")
public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
// @formatter:off
@@ -158,6 +159,9 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions {
private Set<String> functions = new HashSet<>();
// @formatter:on
public OpenAiChatOptions() {
}
public static Builder builder() {
return new Builder();
}

View File

@@ -44,8 +44,6 @@ import java.util.Objects;
@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 }";
public static OpenAiChatResponseMetadata from(OpenAiApi.ChatCompletion result) {
Assert.notNull(result, "OpenAI ChatCompletionResult must not be null");
OpenAiUsage usage = OpenAiUsage.from(result.usage());
@@ -107,7 +105,8 @@ public class OpenAiChatResponseMetadata implements ChatResponseMetadata {
@Override
public String toString() {
return AI_METADATA_STRING.formatted(getClass().getName(), getId(), getUsage(), getRateLimit());
return "OpenAiChatResponseMetadata{" + "id='" + id + '\'' + ", rateLimit=" + rateLimit + ", usage=" + usage
+ ", promptMetadata=" + promptMetadata + '}';
}
@Override

View File

@@ -22,6 +22,7 @@ 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.OpenAiChatOptions;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadata;
import org.springframework.ai.openai.metadata.OpenAiRateLimit;
@@ -31,7 +32,7 @@ import java.time.Duration;
import static org.assertj.core.api.Assertions.assertThat;
public class OpenAiChatResponseTests {
public class OpenAiChatResponseMetadataTests {
@Test
void serDeserChatResponseMetadata() throws JsonProcessingException {
@@ -51,6 +52,21 @@ public class OpenAiChatResponseTests {
assertThat(chatResponseMetadata).usingRecursiveComparison().isEqualTo(deserialized);
}
@Test
void serDeserChatOptions() throws JsonProcessingException {
OpenAiChatOptions openAiChatOptions = OpenAiChatOptions.builder().withModel("mymodel").build();
ObjectMapper objectMapper = new ObjectMapper();
objectMapper.enable(SerializationFeature.INDENT_OUTPUT);
objectMapper.registerModule(new JavaTimeModule());
String json = objectMapper.writeValueAsString(openAiChatOptions);
System.out.println("OpenAiChatOptions Ser: " + json);
OpenAiChatOptions deserialized = objectMapper.readValue(json, OpenAiChatOptions.class);
assertThat(openAiChatOptions).usingRecursiveComparison().isEqualTo(deserialized);
}
@Test
void promptSerialization() throws JsonProcessingException {
Prompt prompt = new Prompt(new UserMessage("hello world"));