From 47c9fcea27c7ba4e26b8e19e16fbe4b94f1ab7ae Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Thu, 23 May 2024 11:16:32 +0200 Subject: [PATCH] Add overload Resourse constructors for fluent API. Fix issue with PromptTemplate failing on trailing parameters not used in the system/user text. --- .../springframework/ai/chat/ChatClient.java | 125 ++++++++++-- .../ai/chat/ChatClientTest.java | 188 +++++++++++++++++- 2 files changed, 292 insertions(+), 21 deletions(-) diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java index 68c381d65..12749b60d 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java @@ -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 extractPlaceholders(String text) { + var placeholders = new ArrayList(); + 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 messages = new ArrayList<>(); - private final Map userParams = new HashMap<>(); + private final Map userParams = new ConcurrentHashMap<>(); - private final Map systemParams = new HashMap<>(); + private final Map 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 functionCallbacks, List functionNames, List media, - ChatOptions chatOptions) { + public ChatClientRequest(ChatModel chatModel, String userText, Map userParams, + String systemText, Map systemParams, List functionCallbacks, + List functionNames, List 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 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 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 pruneTrailingParams(String text, Map params) { + if (CollectionUtils.isEmpty(params)) { + return params; + } + List 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 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 systemSpecConsumer) { this.defaultRequest.system(systemSpecConsumer); return this; diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java index 44645a4e4..02c927769 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java @@ -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 promptCaptor; - @BeforeEach - public void beforeAll() { + private String join(Flux 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"); }