Add ChatOptions to PromptTemplate create methods

Enable ChatOptions configuration when creating Prompts from Templates,

Co-authored-by: zhangqian9158 <zhangqian9158@users.noreply.github.com>
This commit is contained in:
Mark Pollack
2024-08-20 18:14:50 -04:00
parent 528dc0438a
commit dad30f06d9
4 changed files with 78 additions and 0 deletions

View File

@@ -75,10 +75,22 @@ public class ChatPromptTemplate implements PromptTemplateActions, PromptTemplate
return new Prompt(messages);
}
@Override
public Prompt create(ChatOptions modelOptions) {
List<Message> messages = createMessages();
return new Prompt(messages, modelOptions);
}
@Override
public Prompt create(Map<String, Object> model) {
List<Message> messages = createMessages(model);
return new Prompt(messages);
}
@Override
public Prompt create(Map<String, Object> model, ChatOptions modelOptions) {
List<Message> messages = createMessages(model);
return new Prompt(messages, modelOptions);
}
}

View File

@@ -172,11 +172,21 @@ public class PromptTemplate implements PromptTemplateActions, PromptTemplateMess
return new Prompt(render(new HashMap<>()));
}
@Override
public Prompt create(ChatOptions modelOptions) {
return new Prompt(render(new HashMap<>()), modelOptions);
}
@Override
public Prompt create(Map<String, Object> model) {
return new Prompt(render(model));
}
@Override
public Prompt create(Map<String, Object> model, ChatOptions modelOptions) {
return new Prompt(render(model), modelOptions);
}
public Set<String> getInputVariables() {
TokenStream tokens = this.st.impl.tokens;
Set<String> inputVariables = new HashSet<>();

View File

@@ -21,6 +21,10 @@ public interface PromptTemplateActions extends PromptTemplateStringActions {
Prompt create();
Prompt create(ChatOptions modelOptions);
Prompt create(Map<String, Object> model);
Prompt create(Map<String, Object> model, ChatOptions modelOptions);
}

View File

@@ -18,6 +18,9 @@ package org.springframework.ai.prompt;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.chat.prompt.ChatOptionsBuilder;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.core.io.InputStreamResource;
import org.springframework.core.io.Resource;
@@ -31,11 +34,60 @@ import java.util.List;
import java.util.Map;
import static org.assertj.core.api.Assertions.assertThatExceptionOfType;
import static org.assertj.core.api.AssertionsForClassTypes.assertThat;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertThrows;
public class PromptTemplateTest {
@Test
public void testCreateWithEmptyModelAndChatOptions() {
String template = "This is a test prompt with no variables";
PromptTemplate promptTemplate = new PromptTemplate(template);
ChatOptions chatOptions = ChatOptionsBuilder.builder().withTemperature(0.7f).withTopK(3).build();
Prompt prompt = promptTemplate.create(chatOptions);
assertThat(prompt).isNotNull();
assertThat(prompt.getContents()).isEqualTo(template);
assertThat(prompt.getOptions()).isEqualTo(chatOptions);
}
@Test
public void testCreateWithModelAndChatOptions() {
String template = "Hello, {name}! Your age is {age}.";
Map<String, Object> model = new HashMap<>();
model.put("name", "Alice");
model.put("age", 30);
PromptTemplate promptTemplate = new PromptTemplate(template, model);
ChatOptions chatOptions = ChatOptionsBuilder.builder().withTemperature(0.5f).withMaxTokens(100).build();
Prompt prompt = promptTemplate.create(model, chatOptions);
assertThat(prompt).isNotNull();
assertThat(prompt.getContents()).isEqualTo("Hello, Alice! Your age is 30.");
assertThat(prompt.getOptions()).isEqualTo(chatOptions);
}
@Test
public void testCreateWithOverriddenModelAndChatOptions() {
String template = "Hello, {name}! Your favorite color is {color}.";
Map<String, Object> initialModel = new HashMap<>();
initialModel.put("name", "Bob");
initialModel.put("color", "blue");
PromptTemplate promptTemplate = new PromptTemplate(template, initialModel);
Map<String, Object> overriddenModel = new HashMap<>();
overriddenModel.put("color", "red");
ChatOptions chatOptions = ChatOptionsBuilder.builder().withTemperature(0.8f).build();
Prompt prompt = promptTemplate.create(overriddenModel, chatOptions);
assertThat(prompt).isNotNull();
assertThat(prompt.getContents()).isEqualTo("Hello, Bob! Your favorite color is red.");
assertThat(prompt.getOptions()).isEqualTo(chatOptions);
}
@Test
public void testRenderWithList() {
String templateString = "The items are:\n{items:{item | - {item}\n}}";