From f26a6a45671499d5cc180f493a55bbebb8647a07 Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Thu, 18 Apr 2024 13:09:47 -0400 Subject: [PATCH] Add Node data type abstraction for Document and Message --- .../ai/huggingface/client/ClientIT.java | 4 +-- .../ai/chat/messages/AbstractMessage.java | 28 +++++++++---------- .../ai/chat/messages/AssistantMessage.java | 2 +- .../ai/chat/messages/FunctionMessage.java | 2 +- .../ai/chat/messages/Message.java | 6 ++-- .../ai/chat/messages/SystemMessage.java | 2 +- .../ai/chat/messages/UserMessage.java | 2 +- .../springframework/ai/document/Document.java | 5 +++- .../modules/ROOT/pages/api/chatclient.adoc | 17 +++++++++-- .../antora/modules/ROOT/pages/api/prompt.adoc | 16 +++++++++-- 10 files changed, 55 insertions(+), 29 deletions(-) diff --git a/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/client/ClientIT.java b/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/client/ClientIT.java index 2fa2e89bb..f21c7b9ab 100644 --- a/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/client/ClientIT.java +++ b/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/client/ClientIT.java @@ -57,8 +57,8 @@ public class ClientIT { } ```"""; assertThat(chatResponse.getResult().getOutput().getContent()).isEqualTo(expectedResponse); - assertThat(chatResponse.getResult().getOutput().getProperties()).containsKey("generated_tokens"); - assertThat(chatResponse.getResult().getOutput().getProperties()).containsEntry("generated_tokens", 39); + assertThat(chatResponse.getResult().getOutput().getMetadata()).containsKey("generated_tokens"); + assertThat(chatResponse.getResult().getOutput().getMetadata()).containsEntry("generated_tokens", 39); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AbstractMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AbstractMessage.java index 77b54afea..ff78138ca 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AbstractMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AbstractMessage.java @@ -29,7 +29,7 @@ import org.springframework.util.StreamUtils; /** * The AbstractMessage class is an abstract implementation of the Message interface. It - * provides a base implementation for message content, media attachments, properties, and + * provides a base implementation for message content, media attachments, metadata, and * message type. * * @see Message @@ -45,18 +45,18 @@ public abstract class AbstractMessage implements Message { /** * Additional options for the message to influence the response, not a generative map. */ - protected final Map properties; + protected final Map metadata; protected AbstractMessage(MessageType messageType, String content) { this(messageType, content, Map.of()); } - protected AbstractMessage(MessageType messageType, String content, Map messageProperties) { + protected AbstractMessage(MessageType messageType, String content, Map metadata) { Assert.notNull(messageType, "Message type must not be null"); this.messageType = messageType; this.textContent = content; this.mediaData = new ArrayList<>(); - this.properties = messageProperties; + this.metadata = metadata; } protected AbstractMessage(MessageType messageType, String textContent, List mediaData) { @@ -64,7 +64,7 @@ public abstract class AbstractMessage implements Message { } protected AbstractMessage(MessageType messageType, String textContent, List mediaData, - Map messageProperties) { + Map metadata) { Assert.notNull(messageType, "Message type must not be null"); Assert.notNull(textContent, "Content must not be null"); @@ -73,7 +73,7 @@ public abstract class AbstractMessage implements Message { this.messageType = messageType; this.textContent = textContent; this.mediaData = new ArrayList<>(mediaData); - this.properties = messageProperties; + this.metadata = metadata; } protected AbstractMessage(MessageType messageType, Resource resource) { @@ -81,12 +81,12 @@ public abstract class AbstractMessage implements Message { } @SuppressWarnings("null") - protected AbstractMessage(MessageType messageType, Resource resource, Map messageProperties) { + protected AbstractMessage(MessageType messageType, Resource resource, Map metadata) { Assert.notNull(messageType, "Message type must not be null"); Assert.notNull(resource, "Resource must not be null"); this.messageType = messageType; - this.properties = messageProperties; + this.metadata = metadata; this.mediaData = new ArrayList<>(); try (InputStream inputStream = resource.getInputStream()) { @@ -108,8 +108,8 @@ public abstract class AbstractMessage implements Message { } @Override - public Map getProperties() { - return this.properties; + public Map getMetadata() { + return this.metadata; } @Override @@ -122,7 +122,7 @@ public abstract class AbstractMessage implements Message { final int prime = 31; int result = 1; result = prime * result + ((mediaData == null) ? 0 : mediaData.hashCode()); - result = prime * result + ((properties == null) ? 0 : properties.hashCode()); + result = prime * result + ((metadata == null) ? 0 : metadata.hashCode()); result = prime * result + ((messageType == null) ? 0 : messageType.hashCode()); return result; } @@ -142,11 +142,11 @@ public abstract class AbstractMessage implements Message { } else if (!mediaData.equals(other.mediaData)) return false; - if (properties == null) { - if (other.properties != null) + if (metadata == null) { + if (other.metadata != null) return false; } - else if (!properties.equals(other.properties)) + else if (!metadata.equals(other.metadata)) return false; if (messageType != other.messageType) return false; diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AssistantMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AssistantMessage.java index c5f6831ab..f8b890416 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AssistantMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AssistantMessage.java @@ -35,7 +35,7 @@ public class AssistantMessage extends AbstractMessage { @Override public String toString() { - return "AssistantMessage{" + "content='" + getContent() + '\'' + ", properties=" + properties + ", messageType=" + return "AssistantMessage{" + "content='" + getContent() + '\'' + ", properties=" + metadata + ", messageType=" + messageType + '}'; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/FunctionMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/FunctionMessage.java index 06485ac57..a05ef5226 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/FunctionMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/FunctionMessage.java @@ -33,7 +33,7 @@ public class FunctionMessage extends AbstractMessage { @Override public String toString() { - return "FunctionMessage{" + "content='" + getContent() + '\'' + ", properties=" + properties + ", messageType=" + return "FunctionMessage{" + "content='" + getContent() + '\'' + ", properties=" + metadata + ", messageType=" + messageType + '}'; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/Message.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/Message.java index 77e3b5aba..e25f1271d 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/Message.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/Message.java @@ -15,6 +15,8 @@ */ package org.springframework.ai.chat.messages; +import org.springframework.ai.node.Node; + import java.util.List; import java.util.Map; @@ -26,14 +28,12 @@ import java.util.Map; * @see Media * @see MessageType */ -public interface Message { +public interface Message extends Node { String getContent(); List getMedia(); - Map getProperties(); - MessageType getMessageType(); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/SystemMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/SystemMessage.java index 0c0a20e04..8a9dc5eaa 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/SystemMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/SystemMessage.java @@ -36,7 +36,7 @@ public class SystemMessage extends AbstractMessage { @Override public String toString() { - return "SystemMessage{" + "content='" + getContent() + '\'' + ", properties=" + properties + ", messageType=" + return "SystemMessage{" + "content='" + getContent() + '\'' + ", properties=" + metadata + ", messageType=" + messageType + '}'; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/UserMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/UserMessage.java index 4c9229516..6f101579e 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/UserMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/UserMessage.java @@ -45,7 +45,7 @@ public class UserMessage extends AbstractMessage { @Override public String toString() { - return "UserMessage{" + "content='" + getContent() + '\'' + ", properties=" + properties + ", messageType=" + return "UserMessage{" + "content='" + getContent() + '\'' + ", properties=" + metadata + ", messageType=" + messageType + '}'; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/document/Document.java b/spring-ai-core/src/main/java/org/springframework/ai/document/Document.java index 7c54b4fec..8727024eb 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/document/Document.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/document/Document.java @@ -27,6 +27,7 @@ import com.fasterxml.jackson.annotation.JsonProperty; import org.springframework.ai.document.id.IdGenerator; import org.springframework.ai.document.id.RandomIdGenerator; +import org.springframework.ai.node.Node; import org.springframework.util.Assert; /** @@ -34,7 +35,7 @@ import org.springframework.util.Assert; * the document's unique ID and an optional embedding. */ @JsonIgnoreProperties({ "contentFormatter" }) -public class Document { +public class Document implements Node { public final static ContentFormatter DEFAULT_CONTENT_FORMATTER = DefaultContentFormatter.defaultConfig(); @@ -93,6 +94,7 @@ public class Document { return id; } + @Override public String getContent() { return this.content; } @@ -129,6 +131,7 @@ public class Document { this.contentFormatter = contentFormatter; } + @Override public Map getMetadata() { return this.metadata; } diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chatclient.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chatclient.adoc index 902d0d59b..42f7857a5 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chatclient.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chatclient.adoc @@ -78,16 +78,29 @@ The `Message` interface encapsulates a textual message, a collection of attribut [source,java] ---- -public interface Message { +public interface Message extends Node { String getContent(); - Map getProperties(); + List getMedia(); MessageType getMessageType(); } ---- + +and the Node interface is + +```java + +public interface Node { + + T getContent(); + + Map getMetadata(); +} +``` + The `Message` interface has various implementations that correspond to the categories of messages that an AI model can process. Some models, like OpenAI's chat completion endpoint, distinguish between message categories based on conversational roles, effectively mapped by the `MessageType`. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/prompt.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/prompt.adoc index 054d829bd..c568ac4f9 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/prompt.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/prompt.adoc @@ -49,19 +49,29 @@ The `Message` interface encapsulates a textual message, a collection of attribut The interface is defined as follows: ```java -public interface Message { +public interface Message extends Node { String getContent(); List getMedia(); - Map getProperties(); - MessageType getMessageType(); } ``` +and the Node interface is + +```java + +public interface Node { + + T getContent(); + + Map getMetadata(); +} +``` + Various implementations of the `Message` interface correspond to different categories of messages that an AI model can process. Some models, like those from OpenAI, distinguish between message categories based on conversational roles. These roles are effectively mapped by the `MessageType`, as discussed below.