From 9ffc315e1f671a09d5496fa11053458aa262c259 Mon Sep 17 00:00:00 2001 From: Josh Long Date: Sun, 19 May 2024 11:01:05 +0200 Subject: [PATCH] stop the pointless copying of data around and put everything in one single well-known, mutable, carrier object. --- .../springframework/ai/chat/ChatClient.java | 125 ++++++------------ .../ai/chat/DefaultChatClient.java | 21 +-- 2 files changed, 45 insertions(+), 101 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 56e87b805..07925f3ed 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 @@ -31,7 +31,6 @@ import org.springframework.core.io.Resource; import org.springframework.util.Assert; import org.springframework.util.MimeType; import org.springframework.util.StringUtils; -import reactor.core.publisher.Flux; import java.io.IOException; import java.net.URL; @@ -180,20 +179,25 @@ public interface ChatClient { private final List functionCallbacks = new ArrayList<>(); - private final Map userParams = new HashMap<>(); - private final List messages = new ArrayList<>(); + private final Map userParams = new HashMap<>(); + private final Map systemParams = new HashMap<>(); - public ChatClientRequest(ModelCall connector, String userText, String systemText, List functionNames, - List media, ChatOptions chatOptions) { + public ChatClientRequest(ModelCall connector, String userText, String systemText, + List functionCallbacks, List functionNames, List media, + ChatOptions chatOptions) { + + this.connector = connector; + this.chatOptions = chatOptions; + this.userText = userText; this.systemText = systemText; - this.connector = connector; + this.functionNames.addAll(functionNames); + this.functionCallbacks.addAll(functionCallbacks); this.media.addAll(media); - this.chatOptions = chatOptions; } public ChatClientRequest messages(Message... messages) { @@ -222,6 +226,11 @@ public interface ChatClient { return this; } + public ChatClientRequest chatOptions(ChatOptions chatOptions) { + this.chatOptions = chatOptions; + return this; + } + public ChatClientRequest system(Consumer consumer) { var ss = new SystemSpec(); consumer.accept(ss); @@ -270,27 +279,20 @@ public interface ChatClient { } private ChatResponse doGetChatResponse(String processedUserText) { - var messages = new ArrayList(); var textsAreValid = (StringUtils.hasText(processedUserText) || StringUtils.hasText(this.request.systemText)); var messagesAreValid = !this.request.messages.isEmpty(); - Assert.state(!(messagesAreValid && textsAreValid), "you must specify either " + Message.class.getName() + " instances or user/system texts, but not both"); - if (textsAreValid) { - var userMessage = new UserMessage( new PromptTemplate(processedUserText, this.request.userParams).render(), this.request.media); - var systemMessage = new SystemMessage( new PromptTemplate(this.request.systemText, this.request.systemParams).render()); - messages.add(systemMessage); messages.add(userMessage); - } else { messages.addAll(this.request.messages); @@ -311,34 +313,18 @@ public interface ChatClient { return doGetChatResponse(this.request.userText); } - public Flux stream(Class t) { - notSupported(); - return null; - } - - public Flux stream(ParameterizedTypeReference t) { - notSupported(); - return Flux.empty(); - } - public String content() { return doGetChatResponse(this.request.userText).getResult().getOutput().getContent(); } + @SuppressWarnings("unused") public Collection list(Class clzz) { - // todo move to the new ParameterizedTypeReference ready - // BeanOutputConverter - notSupported(); - return null; + return single(new ParameterizedTypeReference>() { + }); } - public Collection list(ParameterizedTypeReference> ptr) { - notSupported(); - return List.of(); - } - - private static void notSupported() { - throw new RuntimeException("this operation is not supported"); + public Collection list(ParameterizedTypeReference> ptr) { + return single(ptr); } } @@ -353,71 +339,42 @@ public interface ChatClient { private final ModelCall modelCall; - private final List defaultMedia = new ArrayList<>(); - - private final List defaultFunctionsNames = new ArrayList<>(); - - private final List defaultFunctionCallbacks = new ArrayList<>(); - - private String defaultSystem; - - private String defaultUser; + private final ChatClientRequest defaultRequest; ChatClientBuilder(ModelCall modelCall) { Assert.notNull(modelCall, "the " + ModelCall.class.getName() + " must be non-null"); this.modelCall = modelCall; + this.defaultRequest = new ChatClientRequest(modelCall, "", "", List.of(), List.of(), List.of(), null); } public ChatClient build() { - return new DefaultChatClient(this.modelCall, this.defaultSystem, this.defaultUser, - this.defaultFunctionsNames, this.defaultMedia); + return new DefaultChatClient(this.modelCall, this.defaultRequest); } - public ChatClientBuilder defaultSystem(Resource resource) { - return this.defaultSystem(resource, Charset.defaultCharset()); - } - - public ChatClientBuilder defaultSystem(Resource resource, Charset charset) { - try { - this.defaultSystem = resource.getContentAsString(charset); - } - catch (IOException e) { - throw new RuntimeException(e); - } + public ChatClientBuilder defaultChatOptions(ChatOptions chatOptions) { + this.defaultRequest.chatOptions(chatOptions); return this; } - public ChatClientBuilder defaultSystem(String systemText) { - this.defaultSystem = systemText; + public ChatClientBuilder defaultUser(Consumer userSpecConsumer) { + this.defaultRequest.user(userSpecConsumer); + return this; + } + + public ChatClientBuilder defaultSystem(Consumer systemSpecConsumer) { + this.defaultRequest.system(systemSpecConsumer); + return this; + } + + public ChatClientBuilder defaultFunctionWrappers(String name, String description, + java.util.function.Function function) { + + this.defaultRequest.function(name, description, function); return this; } public ChatClientBuilder defaultFunctions(String... functionNames) { - this.defaultFunctionsNames.addAll(List.of(functionNames)); - return this; - } - - public ChatClientBuilder defaultFunctions(List functions) { - this.defaultFunctionCallbacks.addAll(functions); - return this; - } - - public ChatClientBuilder defaultUser(Resource userText) { - return this.defaultUser(userText, Charset.defaultCharset()); - } - - public ChatClientBuilder defaultUser(Resource userText, Charset charset) { - try { - this.defaultUser = userText.getContentAsString(charset); - } - catch (IOException e) { - throw new RuntimeException(e); - } - return this; - } - - public ChatClientBuilder defaultUser(String userText) { - this.defaultUser = userText; + this.defaultRequest.functions(functionNames); return this; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/DefaultChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/DefaultChatClient.java index 698049ecb..a09eb1436 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/DefaultChatClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/DefaultChatClient.java @@ -1,10 +1,7 @@ package org.springframework.ai.chat; -import org.springframework.ai.chat.messages.Media; import org.springframework.ai.chat.prompt.Prompt; -import java.util.List; - /** * @author Mark Pollack * @author Christian Tzolov @@ -15,26 +12,16 @@ class DefaultChatClient implements ChatClient { private final ModelCall modelCall; - private final String userText, systemText; + private final ChatClientRequest defaultChatClientRequest; - private final List functionNames; - - private final List media; - - public DefaultChatClient(ModelCall modelCall, String defaultSystemPrompt, String defaultUserPrompt, - List defaultFunctions, List defaultMedia) { + public DefaultChatClient(ModelCall modelCall, ChatClientRequest defaultChatClientRequest) { this.modelCall = modelCall; - this.userText = defaultUserPrompt; - this.systemText = defaultSystemPrompt; - this.functionNames = defaultFunctions; - this.media = defaultMedia; - + this.defaultChatClientRequest = defaultChatClientRequest; } @Override public ChatClientRequest call() { - return new ChatClientRequest(this.modelCall, this.userText, this.systemText, this.functionNames, this.media, - null); + return this.defaultChatClientRequest; } /**