Add Node data type abstraction for Document and Message

This commit is contained in:
Mark Pollack
2024-04-18 13:09:47 -04:00
committed by Christian Tzolov
parent a5923f5a79
commit f26a6a4567
10 changed files with 55 additions and 29 deletions

View File

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

View File

@@ -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<String, Object> properties;
protected final Map<String, Object> metadata;
protected AbstractMessage(MessageType messageType, String content) {
this(messageType, content, Map.of());
}
protected AbstractMessage(MessageType messageType, String content, Map<String, Object> messageProperties) {
protected AbstractMessage(MessageType messageType, String content, Map<String, Object> 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<Media> mediaData) {
@@ -64,7 +64,7 @@ public abstract class AbstractMessage implements Message {
}
protected AbstractMessage(MessageType messageType, String textContent, List<Media> mediaData,
Map<String, Object> messageProperties) {
Map<String, Object> 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<String, Object> messageProperties) {
protected AbstractMessage(MessageType messageType, Resource resource, Map<String, Object> 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<String, Object> getProperties() {
return this.properties;
public Map<String, Object> 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;

View File

@@ -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 + '}';
}

View File

@@ -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 + '}';
}

View File

@@ -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> {
String getContent();
List<Media> getMedia();
Map<String, Object> getProperties();
MessageType getMessageType();
}

View File

@@ -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 + '}';
}

View File

@@ -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 + '}';
}

View File

@@ -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<String> {
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<String, Object> getMetadata() {
return this.metadata;
}

View File

@@ -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> {
String getContent();
Map<String, Object> getProperties();
List<Media> getMedia();
MessageType getMessageType();
}
----
and the Node interface is
```java
public interface Node<T> {
T getContent();
Map<String, Object> 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`.

View File

@@ -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> {
String getContent();
List<Media> getMedia();
Map<String, Object> getProperties();
MessageType getMessageType();
}
```
and the Node interface is
```java
public interface Node<T> {
T getContent();
Map<String, Object> 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.