Add overload Resourse constructors for fluent API.

Fix issue with PromptTemplate failing on trailing parameters not used in the system/user text.
This commit is contained in:
Christian Tzolov
2024-05-23 11:16:32 +02:00
parent 5add8f0b68
commit 47c9fcea27
2 changed files with 292 additions and 21 deletions

View File

@@ -24,7 +24,10 @@ import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.function.Consumer;
import java.util.regex.Pattern;
import java.util.stream.Collectors;
import reactor.core.publisher.Flux;
@@ -57,6 +60,17 @@ import org.springframework.util.StringUtils;
*/
public interface ChatClient {
Pattern PLACEHOLDER_EXTRACTION_PATTER = Pattern.compile("\\{(.*?)\\}");
private static List<String> extractPlaceholders(String text) {
var placeholders = new ArrayList<String>();
var matcher = PLACEHOLDER_EXTRACTION_PATTER.matcher(text);
while (matcher.find()) {
placeholders.add(matcher.group(1));
}
return placeholders;
}
// static ChatClient create(ChatModel chatModel) {
// return builder(chatModel).build();
// }
@@ -214,25 +228,27 @@ public interface ChatClient {
private final List<Message> messages = new ArrayList<>();
private final Map<String, Object> userParams = new HashMap<>();
private final Map<String, Object> userParams = new ConcurrentHashMap<>();
private final Map<String, Object> systemParams = new HashMap<>();
private final Map<String, Object> systemParams = new ConcurrentHashMap<>();
/* copy constructor */
ChatClientRequest(ChatClientRequest ccr) {
this(ccr.chatModel, ccr.userText, ccr.systemText, ccr.functionCallbacks, ccr.functionNames, ccr.media,
ccr.chatOptions);
this(ccr.chatModel, ccr.userText, ccr.userParams, ccr.systemText, ccr.systemParams, ccr.functionCallbacks,
ccr.functionNames, ccr.media, ccr.chatOptions);
}
public ChatClientRequest(ChatModel chatModel, String userText, String systemText,
List<FunctionCallback> functionCallbacks, List<String> functionNames, List<Media> media,
ChatOptions chatOptions) {
public ChatClientRequest(ChatModel chatModel, String userText, Map<String, Object> userParams,
String systemText, Map<String, Object> systemParams, List<FunctionCallback> functionCallbacks,
List<String> functionNames, List<Media> media, ChatOptions chatOptions) {
this.chatModel = chatModel;
this.chatOptions = chatOptions != null ? chatOptions : chatModel.getDefaultOptions();
this.userText = userText;
this.userParams.putAll(userParams);
this.systemText = systemText;
this.systemParams.putAll(systemParams);
this.functionNames.addAll(functionNames);
this.functionCallbacks.addAll(functionCallbacks);
@@ -280,11 +296,26 @@ public interface ChatClient {
return this;
}
public ChatClientRequest system(Resource text, Charset charset) {
try {
this.systemText = text.getContentAsString(charset);
}
catch (IOException e) {
throw new RuntimeException(e);
}
return this;
}
public ChatClientRequest system(Resource text) {
return this.system(text, Charset.defaultCharset());
}
public ChatClientRequest system(Consumer<SystemSpec> consumer) {
var ss = new SystemSpec();
consumer.accept(ss);
this.systemText = StringUtils.hasText(ss.text()) ? ss.text() : this.systemText;
this.systemParams.putAll(ss.params());
return this;
}
@@ -293,6 +324,20 @@ public interface ChatClient {
return this;
}
public ChatClientRequest user(Resource text, Charset charset) {
try {
this.userText = text.getContentAsString(charset);
}
catch (IOException e) {
throw new RuntimeException(e);
}
return this;
}
public ChatClientRequest user(Resource text) {
return this.user(text, Charset.defaultCharset());
}
public ChatClientRequest user(Consumer<UserSpec> consumer) {
var us = new UserSpec();
consumer.accept(us);
@@ -365,6 +410,19 @@ public interface ChatClient {
}
// Hack: Prune any trailing parameters not used in the system text.
// Later will cause the ST string template to fail.
private static Map<String, Object> pruneTrailingParams(String text, Map<String, Object> params) {
if (CollectionUtils.isEmpty(params)) {
return params;
}
List<String> paramNames = extractPlaceholders(text);
return params.entrySet()
.stream()
.filter(e -> paramNames.contains(e.getKey()))
.collect(Collectors.toMap(e -> e.getKey(), e -> e.getValue()));
}
public static class CallResponseSpec {
private final ChatClientRequest request;
@@ -413,15 +471,19 @@ public interface ChatClient {
if (textsAreValid) {
UserMessage userMessage = null;
if (!CollectionUtils.isEmpty(userParams)) {
userMessage = new UserMessage(new PromptTemplate(processedUserText, userParams).render(),
userMessage = new UserMessage(
new PromptTemplate(processedUserText,
pruneTrailingParams(processedUserText, userParams))
.render(),
this.request.media);
}
else {
userMessage = new UserMessage(processedUserText, this.request.media);
}
if (StringUtils.hasText(this.request.systemText) || !this.request.systemParams.isEmpty()) {
var systemMessage = new SystemMessage(
new PromptTemplate(this.request.systemText, this.request.systemParams).render());
var systemMessage = new SystemMessage(new PromptTemplate(this.request.systemText,
pruneTrailingParams(this.request.systemText, this.request.systemParams))
.render());
messages.add(systemMessage);
}
messages.add(userMessage);
@@ -484,15 +546,19 @@ public interface ChatClient {
if (textsAreValid) {
UserMessage userMessage = null;
if (!CollectionUtils.isEmpty(userParams)) {
userMessage = new UserMessage(new PromptTemplate(processedUserText, userParams).render(),
userMessage = new UserMessage(
new PromptTemplate(processedUserText,
pruneTrailingParams(processedUserText, userParams))
.render(),
this.request.media);
}
else {
userMessage = new UserMessage(processedUserText, this.request.media);
}
if (StringUtils.hasText(this.request.systemText) || !this.request.systemParams.isEmpty()) {
var systemMessage = new SystemMessage(
new PromptTemplate(this.request.systemText, this.request.systemParams).render());
var systemMessage = new SystemMessage(new PromptTemplate(this.request.systemText,
pruneTrailingParams(this.request.systemText, this.request.systemParams))
.render());
messages.add(systemMessage);
}
messages.add(userMessage);
@@ -550,14 +616,15 @@ public interface ChatClient {
ChatClientBuilder(ChatModel chatModel) {
Assert.notNull(chatModel, "the " + ChatModel.class.getName() + " must be non-null");
this.chatModel = chatModel;
this.defaultRequest = new ChatClientRequest(chatModel, "", "", List.of(), List.of(), List.of(), null);
this.defaultRequest = new ChatClientRequest(chatModel, "", Map.of(), "", Map.of(), List.of(), List.of(),
List.of(), null);
}
public ChatClient build() {
return new DefaultChatClient(this.chatModel, this.defaultRequest);
}
public ChatClientBuilder defaultRuntimeOptions(ChatOptions chatOptions) {
public ChatClientBuilder defaultOptions(ChatOptions chatOptions) {
this.defaultRequest.chatOptions(chatOptions);
return this;
}
@@ -567,6 +634,20 @@ public interface ChatClient {
return this;
}
public ChatClientBuilder defaultUser(Resource text, Charset charset) {
try {
this.defaultRequest.user(text.getContentAsString(charset));
}
catch (IOException e) {
throw new RuntimeException(e);
}
return this;
}
public ChatClientBuilder defaultUser(Resource text) {
return this.defaultUser(text, Charset.defaultCharset());
}
public ChatClientBuilder defaultUser(Consumer<UserSpec> userSpecConsumer) {
this.defaultRequest.user(userSpecConsumer);
return this;
@@ -577,6 +658,20 @@ public interface ChatClient {
return this;
}
public ChatClientBuilder defaultSystem(Resource text, Charset charset) {
try {
this.defaultRequest.system(text.getContentAsString(charset));
}
catch (IOException e) {
throw new RuntimeException(e);
}
return this;
}
public ChatClientBuilder defaultSystem(Resource text) {
return this.defaultSystem(text, Charset.defaultCharset());
}
public ChatClientBuilder defaultSystem(Consumer<SystemSpec> systemSpecConsumer) {
this.defaultRequest.system(systemSpecConsumer);
return this;

View File

@@ -19,14 +19,15 @@ package org.springframework.ai.chat;
import java.net.MalformedURLException;
import java.net.URL;
import java.util.List;
import java.util.stream.Collectors;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.ArgumentCaptor;
import org.mockito.Captor;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;
import reactor.core.publisher.Flux;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.MessageType;
@@ -44,22 +45,189 @@ import static org.mockito.Mockito.when;
@ExtendWith(MockitoExtension.class)
public class ChatClientTest {
public static interface MixChatModel extends ChatModel, StreamingChatModel {
}
@Mock
ChatModel chatModel;
MixChatModel chatModel;
@Captor
ArgumentCaptor<Prompt> promptCaptor;
@BeforeEach
public void beforeAll() {
private String join(Flux<String> fluxContent) {
return fluxContent.collectList().block().stream().collect(Collectors.joining());
}
// ChatClient Builder Tests
@Test
public void defaultSystemText() {
when(chatModel.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
when(chatModel.stream(promptCaptor.capture()))
.thenReturn(
Flux.generate(() -> new ChatResponse(List.of(new Generation("response"))), (state, sink) -> {
sink.next(state);
sink.complete();
return state;
}));
var chatClient = ChatClient.builder(chatModel)
.defaultSystem("Default system text").build();
var content = chatClient.prompt().call().content();
assertThat(content).isEqualTo("response");
Message systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Default system text");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
content = join(chatClient.prompt().stream().content());
assertThat(content).isEqualTo("response");
systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Default system text");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
// Override the default system text with prompt system
content = chatClient.prompt()
.system("Override default system text")
.call().content();
assertThat(content).isEqualTo("response");
systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Override default system text");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
// Streaming
content = join(chatClient.prompt()
.system("Override default system text")
.stream().content());
assertThat(content).isEqualTo("response");
systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Override default system text");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
}
@Test
public void defaultSystemTextLambda() {
when(chatModel.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
when(chatModel.stream(promptCaptor.capture()))
.thenReturn(
Flux.generate(() -> new ChatResponse(List.of(new Generation("response"))), (state, sink) -> {
sink.next(state);
sink.complete();
return state;
}));
var chatClient = ChatClient.builder(chatModel)
.defaultSystem(s -> s.text("Default system text {param1}, {param2}")
.param("param1", "value1")
.param("param2", "value2"))
.build();
var content = chatClient.prompt().call().content();
assertThat(content).isEqualTo("response");
Message systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Default system text value1, value2");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
// Streaming
content = join(chatClient.prompt().stream().content());
assertThat(content).isEqualTo("response");
systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Default system text value1, value2");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
// Override single default system parameter
content = chatClient.prompt()
.system(s -> s.param("param1", "value1New"))
.call().content();
assertThat(content).isEqualTo("response");
systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Default system text value1New, value2");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
content = join(chatClient.prompt()
.system(s -> s.param("param1", "value1New"))
.stream().content());
assertThat(content).isEqualTo("response");
systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Default system text value1New, value2");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
// Override default system text
content = chatClient.prompt()
.system(s -> s.text("Override default system text {param3}")
.param("param3", "value3"))
.call().content();
assertThat(content).isEqualTo("response");
systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Override default system text value3");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
// Streaming
content = join(chatClient.prompt()
.system(s -> s.text("Override default system text {param3}")
.param("param3", "value3"))
.stream().content());
assertThat(content).isEqualTo("response");
systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("Override default system text value3");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
}
@Test
public void defaultUserText() {
when(chatModel.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
var chatClient = ChatClient.builder(chatModel)
.defaultUser("Default user text").build();
var content = chatClient.prompt().call().content();
assertThat(content).isEqualTo("response");
Message userMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(userMessage.getContent()).isEqualTo("Default user text");
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
// Override the default system text with prompt system
content = chatClient.prompt()
.user("Override default user text")
.call().content();
assertThat(content).isEqualTo("response");
userMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(userMessage.getContent()).isEqualTo("Override default user text");
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
}
@Test
public void simpleUserPrompt() {
when(chatModel.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
assertThat(ChatClient.builder(chatModel).build().prompt().user("User prompt").call().content())
.isEqualTo("response");
.isEqualTo("response");
Message userMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(userMessage.getContent()).isEqualTo("User prompt");
@@ -68,6 +236,9 @@ public class ChatClientTest {
@Test
public void simpleUserPromptObject() throws MalformedURLException {
when(chatModel.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
UserMessage message = new UserMessage("User prompt");
Prompt prompt = new Prompt(message);
assertThat(ChatClient.builder(chatModel).build().prompt(prompt).call().content()).isEqualTo("response");
@@ -79,6 +250,9 @@ public class ChatClientTest {
@Test
public void simpleSystemPrompt() throws MalformedURLException {
when(chatModel.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
String response = ChatClient.builder(chatModel).build().prompt().system("System prompt").call().content();
assertThat(response).isEqualTo("response");
@@ -97,6 +271,8 @@ public class ChatClientTest {
@Test
public void complexCall() throws MalformedURLException {
when(chatModel.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
var options = FunctionCallingOptions.builder().build();
when(chatModel.getDefaultOptions()).thenReturn(options);
@@ -128,7 +304,7 @@ public class ChatClientTest {
assertThat(userMessage.getMedia()).hasSize(1);
assertThat(userMessage.getMedia().iterator().next().getMimeType()).isEqualTo(MimeTypeUtils.IMAGE_PNG);
assertThat(userMessage.getMedia().iterator().next().getData())
.isEqualTo("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png");
.isEqualTo("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png");
assertThat(options.getFunctions()).containsExactly("function1");
}