Add toString to Message classes, and allow Resource to be used in PromptTemplate

This commit is contained in:
Mark Pollack
2023-08-15 23:22:32 -04:00
parent 811ff1ed2b
commit d4b9d071ba
6 changed files with 70 additions and 7 deletions

View File

@@ -16,12 +16,12 @@
package org.springframework.ai.prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import java.util.Collections;
import java.util.List;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.prompt.messages.Message;
public class Prompt {
private List<Message> messages;
@@ -50,4 +50,9 @@ public class Prompt {
return this.messages;
}
@Override
public String toString() {
return "Prompt{" + "messages=" + messages + '}';
}
}

View File

@@ -20,9 +20,14 @@ import org.antlr.runtime.Token;
import org.antlr.runtime.TokenStream;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.core.io.Resource;
import org.springframework.util.StreamUtils;
import org.stringtemplate.v4.ST;
import org.stringtemplate.v4.compiler.STLexer;
import java.io.IOException;
import java.io.InputStream;
import java.nio.charset.Charset;
import java.util.*;
import java.util.Map.Entry;
import java.util.stream.Collectors;
@@ -40,6 +45,21 @@ public class PromptTemplate implements PromptTemplateActions {
private OutputParser outputParser;
public PromptTemplate(Resource resource) {
try (InputStream inputStream = resource.getInputStream()) {
this.template = StreamUtils.copyToString(inputStream, Charset.defaultCharset());
}
catch (IOException ex) {
throw new RuntimeException("Failed to read resource", ex);
}
try {
this.st = new ST(this.template, '{', '}');
}
catch (Exception ex) {
throw new IllegalArgumentException("The template string is not valid.", ex);
}
}
public PromptTemplate(String template) {
this.template = template;
// If the template string is not valid, an exception will be thrown
@@ -95,12 +115,26 @@ public class PromptTemplate implements PromptTemplateActions {
@Override
public String render(Map<String, Object> model) {
validate(model);
for (Entry<String, Object> stringObjectEntry : model.entrySet()) {
if (st.getAttribute(stringObjectEntry.getKey()) == null) {
st.add(stringObjectEntry.getKey(), stringObjectEntry.getValue());
for (Entry<String, Object> entry : model.entrySet()) {
if (st.getAttribute(entry.getKey()) == null) {
if (entry.getValue() instanceof Resource) {
st.add(entry.getKey(), renderResource((Resource) entry.getValue()));
}
else {
st.add(entry.getKey(), entry.getValue());
}
}
}
return st.render().trim();
return st.render();
}
private String renderResource(Resource resource) {
try (InputStream inputStream = resource.getInputStream()) {
return StreamUtils.copyToString(inputStream, Charset.defaultCharset());
}
catch (IOException ex) {
throw new RuntimeException(ex);
}
}
@Override

View File

@@ -34,4 +34,10 @@ public class AssistantMessage extends AbstractMessage {
super(MessageType.ASSISTANT, content, properties);
}
@Override
public String toString() {
return "AssistantMessage{" + "content='" + content + '\'' + ", properties=" + properties + ", messageType="
+ messageType + '}';
}
}

View File

@@ -28,4 +28,10 @@ public class FunctionMessage extends AbstractMessage {
super(MessageType.SYSTEM, content, properties);
}
@Override
public String toString() {
return "FunctionMessage{" + "content='" + content + '\'' + ", properties=" + properties + ", messageType="
+ messageType + '}';
}
}

View File

@@ -34,4 +34,10 @@ public class SystemMessage extends AbstractMessage {
super(MessageType.SYSTEM, content, properties);
}
@Override
public String toString() {
return "SystemMessage{" + "content='" + content + '\'' + ", properties=" + properties + ", messageType="
+ messageType + '}';
}
}

View File

@@ -33,4 +33,10 @@ public class UserMessage extends AbstractMessage {
super(MessageType.USER, message, properties);
}
@Override
public String toString() {
return "UserMessage{" + "content='" + content + '\'' + ", properties=" + properties + ", messageType="
+ messageType + '}';
}
}