Add (Streaming)ChatClient convinience defaults

Facilitates the creation of multimodal message queries.
This commit is contained in:
Christian Tzolov
2024-04-20 07:05:07 +02:00
parent df92dff36a
commit 2cdd9b5138
3 changed files with 25 additions and 0 deletions

View File

@@ -16,6 +16,10 @@
package org.springframework.ai.chat;
import org.springframework.ai.chat.prompt.Prompt;
import java.util.Arrays;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.model.ModelClient;
@@ -28,6 +32,12 @@ public interface ChatClient extends ModelClient<Prompt, ChatResponse> {
return (generation != null) ? generation.getOutput().getContent() : "";
}
default String call(Message... messages) {
Prompt prompt = new Prompt(Arrays.asList(messages));
Generation generation = call(prompt).getResult();
return (generation != null) ? generation.getOutput().getContent() : "";
}
@Override
ChatResponse call(Prompt prompt);

View File

@@ -15,8 +15,11 @@
*/
package org.springframework.ai.chat;
import java.util.Arrays;
import reactor.core.publisher.Flux;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.model.StreamingModelClient;
@@ -30,6 +33,13 @@ public interface StreamingChatClient extends StreamingModelClient<Prompt, ChatRe
: response.getResult().getOutput().getContent());
}
default Flux<String> call(Message... messages) {
Prompt prompt = new Prompt(Arrays.asList(messages));
return stream(prompt).map(response -> (response.getResult() == null || response.getResult().getOutput() == null
|| response.getResult().getOutput().getContent() == null) ? ""
: response.getResult().getOutput().getContent());
}
@Override
Flux<ChatResponse> stream(Prompt prompt);

View File

@@ -15,6 +15,7 @@
*/
package org.springframework.ai.chat.messages;
import java.util.Arrays;
import java.util.List;
import org.springframework.core.io.Resource;
@@ -38,6 +39,10 @@ public class UserMessage extends AbstractMessage {
super(MessageType.USER, textContent, mediaList);
}
public UserMessage(String textContent, Media... media) {
this(textContent, Arrays.asList(media));
}
@Override
public String toString() {
return "UserMessage{" + "content='" + getContent() + '\'' + ", properties=" + properties + ", messageType="