From 5594cc4befd8280211db6823b5e153e382c439e3 Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Sat, 8 Jun 2024 19:34:33 -0400 Subject: [PATCH] more fixes for openai ser-deser --- .../ai/openai/OpenAiChatOptions.java | 18 +++++++++++------- .../metadata/OpenAiChatResponseMetadata.java | 5 ++--- ...va => OpenAiChatResponseMetadataTests.java} | 18 +++++++++++++++++- 3 files changed, 30 insertions(+), 11 deletions(-) rename models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/{OpenAiChatResponseTests.java => OpenAiChatResponseMetadataTests.java} (78%) diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java index d980c5c78..fe3152311 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java @@ -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 functions = new HashSet<>(); // @formatter:on + public OpenAiChatOptions() { + } + public static Builder builder() { return new Builder(); } 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 71988d659..ff6a728ff 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 @@ -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 diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatResponseTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatResponseMetadataTests.java similarity index 78% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatResponseTests.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatResponseMetadataTests.java index a41d69cda..251fb03dd 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatResponseTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatResponseMetadataTests.java @@ -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"));