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:
@@ -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;
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user