Make ChatClient APIs null-safe and predictable
* Introduced null-safety for all ChatClient APIs and verified via extensive auto-tests. * Made the final message list computed by ChatClient consistent and predictable across the different options for providing messages (prompt(), messages(), user(), prompt()) and verified via extensive auto-tests. * Added even more auto-tests to cover as many scenarios and edge cases as possible. Signed-off-by: Thomas Vitale <ThomasVitale@users.noreply.github.com>
This commit is contained in:
committed by
Christian Tzolov
parent
25f386bb14
commit
048d54b82b
@@ -38,6 +38,8 @@ import org.springframework.ai.model.Media;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.core.ParameterizedTypeReference;
|
||||
import org.springframework.core.io.Resource;
|
||||
import org.springframework.lang.Nullable;
|
||||
import org.springframework.util.Assert;
|
||||
import org.springframework.util.MimeType;
|
||||
|
||||
/**
|
||||
@@ -49,6 +51,7 @@ import org.springframework.util.MimeType;
|
||||
* @author Christian Tzolov
|
||||
* @author Josh Long
|
||||
* @author Arjen Poutsma
|
||||
* @author Thomas Vitale
|
||||
* @since 1.0.0
|
||||
*/
|
||||
public interface ChatClient {
|
||||
@@ -62,7 +65,9 @@ public interface ChatClient {
|
||||
}
|
||||
|
||||
static ChatClient create(ChatModel chatModel, ObservationRegistry observationRegistry,
|
||||
ChatClientObservationConvention observationConvention) {
|
||||
@Nullable ChatClientObservationConvention observationConvention) {
|
||||
Assert.notNull(chatModel, "chatModel cannot be null");
|
||||
Assert.notNull(observationRegistry, "observationRegistry cannot be null");
|
||||
return builder(chatModel, observationRegistry, observationConvention).build();
|
||||
}
|
||||
|
||||
@@ -71,7 +76,9 @@ public interface ChatClient {
|
||||
}
|
||||
|
||||
static Builder builder(ChatModel chatModel, ObservationRegistry observationRegistry,
|
||||
ChatClientObservationConvention customObservationConvention) {
|
||||
@Nullable ChatClientObservationConvention customObservationConvention) {
|
||||
Assert.notNull(chatModel, "chatModel cannot be null");
|
||||
Assert.notNull(observationRegistry, "observationRegistry cannot be null");
|
||||
return new DefaultChatClientBuilder(chatModel, observationRegistry, customObservationConvention);
|
||||
}
|
||||
|
||||
@@ -136,14 +143,19 @@ public interface ChatClient {
|
||||
|
||||
interface CallResponseSpec {
|
||||
|
||||
@Nullable
|
||||
<T> T entity(ParameterizedTypeReference<T> type);
|
||||
|
||||
@Nullable
|
||||
<T> T entity(StructuredOutputConverter<T> structuredOutputConverter);
|
||||
|
||||
@Nullable
|
||||
<T> T entity(Class<T> type);
|
||||
|
||||
@Nullable
|
||||
ChatResponse chatResponse();
|
||||
|
||||
@Nullable
|
||||
String content();
|
||||
|
||||
<T> ResponseEntity<ChatResponse, T> responseEntity(Class<T> type);
|
||||
|
||||
@@ -21,6 +21,7 @@ import java.net.URL;
|
||||
import java.nio.charset.Charset;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Arrays;
|
||||
import java.util.Collection;
|
||||
import java.util.Collections;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
@@ -63,6 +64,7 @@ import org.springframework.ai.model.function.FunctionCallbackWrapper;
|
||||
import org.springframework.core.Ordered;
|
||||
import org.springframework.core.ParameterizedTypeReference;
|
||||
import org.springframework.core.io.Resource;
|
||||
import org.springframework.lang.Nullable;
|
||||
import org.springframework.util.Assert;
|
||||
import org.springframework.util.CollectionUtils;
|
||||
import org.springframework.util.MimeType;
|
||||
@@ -88,18 +90,44 @@ public class DefaultChatClient implements ChatClient {
|
||||
private final DefaultChatClientRequestSpec defaultChatClientRequest;
|
||||
|
||||
public DefaultChatClient(DefaultChatClientRequestSpec defaultChatClientRequest) {
|
||||
Assert.notNull(defaultChatClientRequest, "defaultChatClientRequest cannot be null");
|
||||
this.defaultChatClientRequest = defaultChatClientRequest;
|
||||
}
|
||||
|
||||
private static AdvisedRequest toAdvisedRequest(DefaultChatClientRequestSpec inputRequest, String formatParam) {
|
||||
private static AdvisedRequest toAdvisedRequest(DefaultChatClientRequestSpec inputRequest,
|
||||
@Nullable String formatParam) {
|
||||
Assert.notNull(inputRequest, "inputRequest cannot be null");
|
||||
|
||||
Map<String, Object> advisorContext = new ConcurrentHashMap<>(inputRequest.getAdvisorParams());
|
||||
if (StringUtils.hasText(formatParam)) {
|
||||
advisorContext.put("formatParam", formatParam);
|
||||
}
|
||||
|
||||
return new AdvisedRequest(inputRequest.chatModel, inputRequest.userText, inputRequest.systemText,
|
||||
inputRequest.chatOptions, inputRequest.media, inputRequest.functionNames,
|
||||
inputRequest.functionCallbacks, inputRequest.messages, inputRequest.userParams,
|
||||
// Process userText, media and messages before creating the AdvisedRequest.
|
||||
String userText = inputRequest.userText;
|
||||
List<Media> media = inputRequest.media;
|
||||
List<Message> messages = inputRequest.messages;
|
||||
|
||||
// If the userText is empty, then try extracting the userText from the last
|
||||
// message
|
||||
// in the messages list and remove it from the messages list.
|
||||
if (!StringUtils.hasText(userText) && !CollectionUtils.isEmpty(messages)) {
|
||||
Message lastMessage = messages.get(messages.size() - 1);
|
||||
if (lastMessage.getMessageType() == MessageType.USER) {
|
||||
UserMessage userMessage = (UserMessage) lastMessage;
|
||||
if (StringUtils.hasText(userMessage.getContent())) {
|
||||
userText = lastMessage.getContent();
|
||||
}
|
||||
Collection<Media> messageMedia = userMessage.getMedia();
|
||||
if (!CollectionUtils.isEmpty(messageMedia)) {
|
||||
media.addAll(messageMedia);
|
||||
}
|
||||
messages = messages.subList(0, messages.size() - 1);
|
||||
}
|
||||
}
|
||||
|
||||
return new AdvisedRequest(inputRequest.chatModel, userText, inputRequest.systemText, inputRequest.chatOptions,
|
||||
media, inputRequest.functionNames, inputRequest.functionCallbacks, messages, inputRequest.userParams,
|
||||
inputRequest.systemParams, inputRequest.advisors, inputRequest.advisorParams, advisorContext,
|
||||
inputRequest.toolContext);
|
||||
}
|
||||
@@ -122,10 +150,13 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
@Override
|
||||
public ChatClientRequestSpec prompt(String content) {
|
||||
Assert.hasText(content, "content cannot be null or empty");
|
||||
return prompt(new Prompt(content));
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientRequestSpec prompt(Prompt prompt) {
|
||||
Assert.notNull(prompt, "prompt cannot be null");
|
||||
|
||||
DefaultChatClientRequestSpec spec = new DefaultChatClientRequestSpec(this.defaultChatClientRequest);
|
||||
|
||||
@@ -135,26 +166,10 @@ public class DefaultChatClient implements ChatClient {
|
||||
}
|
||||
|
||||
// Messages
|
||||
List<Message> messages = prompt.getInstructions();
|
||||
|
||||
if (!CollectionUtils.isEmpty(messages)) {
|
||||
var lastMessage = messages.get(messages.size() - 1);
|
||||
if (lastMessage.getMessageType() == MessageType.USER) {
|
||||
// Unzip the last message
|
||||
var userMessage = (UserMessage) lastMessage;
|
||||
if (StringUtils.hasText(userMessage.getContent())) {
|
||||
spec.user(lastMessage.getContent());
|
||||
}
|
||||
var media = userMessage.getMedia();
|
||||
if (!CollectionUtils.isEmpty(media)) {
|
||||
spec.user(u -> u.media(media.toArray(new Media[media.size()])));
|
||||
}
|
||||
messages = messages.subList(0, messages.size() - 1);
|
||||
}
|
||||
if (prompt.getInstructions() != null) {
|
||||
spec.messages(prompt.getInstructions());
|
||||
}
|
||||
|
||||
spec.messages(messages);
|
||||
|
||||
return spec;
|
||||
}
|
||||
|
||||
@@ -173,34 +188,44 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
private final List<Media> media = new ArrayList<>();
|
||||
|
||||
private String text = "";
|
||||
@Nullable
|
||||
private String text;
|
||||
|
||||
@Override
|
||||
public PromptUserSpec media(Media... media) {
|
||||
Assert.notNull(media, "media cannot be null");
|
||||
Assert.noNullElements(media, "media cannot contain null elements");
|
||||
this.media.addAll(Arrays.asList(media));
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptUserSpec media(MimeType mimeType, URL url) {
|
||||
Assert.notNull(mimeType, "mimeType cannot be null");
|
||||
Assert.notNull(url, "url cannot be null");
|
||||
this.media.add(new Media(mimeType, url));
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptUserSpec media(MimeType mimeType, Resource resource) {
|
||||
Assert.notNull(mimeType, "mimeType cannot be null");
|
||||
Assert.notNull(resource, "resource cannot be null");
|
||||
this.media.add(new Media(mimeType, resource));
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptUserSpec text(String text) {
|
||||
Assert.hasText(text, "text cannot be null or empty");
|
||||
this.text = text;
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptUserSpec text(Resource text, Charset charset) {
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
Assert.notNull(charset, "charset cannot be null");
|
||||
try {
|
||||
this.text(text.getContentAsString(charset));
|
||||
}
|
||||
@@ -212,22 +237,29 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
@Override
|
||||
public PromptUserSpec text(Resource text) {
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
this.text(text, Charset.defaultCharset());
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptUserSpec param(String k, Object v) {
|
||||
this.params.put(k, v);
|
||||
public PromptUserSpec param(String key, Object value) {
|
||||
Assert.hasText(key, "key cannot be null or empty");
|
||||
Assert.notNull(value, "value cannot be null");
|
||||
this.params.put(key, value);
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptUserSpec params(Map<String, Object> p) {
|
||||
this.params.putAll(p);
|
||||
public PromptUserSpec params(Map<String, Object> params) {
|
||||
Assert.notNull(params, "params cannot be null");
|
||||
Assert.noNullElements(params.keySet(), "param keys cannot contain null elements");
|
||||
Assert.noNullElements(params.values(), "param values cannot contain null elements");
|
||||
this.params.putAll(params);
|
||||
return this;
|
||||
}
|
||||
|
||||
@Nullable
|
||||
protected String text() {
|
||||
return this.text;
|
||||
}
|
||||
@@ -246,16 +278,20 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
private final Map<String, Object> params = new HashMap<>();
|
||||
|
||||
private String text = "";
|
||||
@Nullable
|
||||
private String text;
|
||||
|
||||
@Override
|
||||
public PromptSystemSpec text(String text) {
|
||||
Assert.hasText(text, "text cannot be null or empty");
|
||||
this.text = text;
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptSystemSpec text(Resource text, Charset charset) {
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
Assert.notNull(charset, "charset cannot be null");
|
||||
try {
|
||||
this.text(text.getContentAsString(charset));
|
||||
}
|
||||
@@ -267,22 +303,29 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
@Override
|
||||
public PromptSystemSpec text(Resource text) {
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
this.text(text, Charset.defaultCharset());
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptSystemSpec param(String k, Object v) {
|
||||
this.params.put(k, v);
|
||||
public PromptSystemSpec param(String key, Object value) {
|
||||
Assert.hasText(key, "key cannot be null or empty");
|
||||
Assert.notNull(value, "value cannot be null");
|
||||
this.params.put(key, value);
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PromptSystemSpec params(Map<String, Object> p) {
|
||||
this.params.putAll(p);
|
||||
public PromptSystemSpec params(Map<String, Object> params) {
|
||||
Assert.notNull(params, "params cannot be null");
|
||||
Assert.noNullElements(params.keySet(), "param keys cannot contain null elements");
|
||||
Assert.noNullElements(params.values(), "param values cannot contain null elements");
|
||||
this.params.putAll(params);
|
||||
return this;
|
||||
}
|
||||
|
||||
@Nullable
|
||||
protected String text() {
|
||||
return this.text;
|
||||
}
|
||||
@@ -299,22 +342,35 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
private final Map<String, Object> params = new HashMap<>();
|
||||
|
||||
public AdvisorSpec param(String k, Object v) {
|
||||
this.params.put(k, v);
|
||||
@Override
|
||||
public AdvisorSpec param(String key, Object value) {
|
||||
Assert.hasText(key, "key cannot be null or empty");
|
||||
Assert.notNull(value, "value cannot be null");
|
||||
this.params.put(key, value);
|
||||
return this;
|
||||
}
|
||||
|
||||
public AdvisorSpec params(Map<String, Object> p) {
|
||||
this.params.putAll(p);
|
||||
@Override
|
||||
public AdvisorSpec params(Map<String, Object> params) {
|
||||
Assert.notNull(params, "params cannot be null");
|
||||
Assert.noNullElements(params.keySet(), "param keys cannot contain null elements");
|
||||
Assert.noNullElements(params.values(), "param values cannot contain null elements");
|
||||
this.params.putAll(params);
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public AdvisorSpec advisors(Advisor... advisors) {
|
||||
Assert.notNull(advisors, "advisors cannot be null");
|
||||
Assert.noNullElements(advisors, "advisors cannot contain null elements");
|
||||
this.advisors.addAll(List.of(advisors));
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public AdvisorSpec advisors(List<Advisor> advisors) {
|
||||
Assert.notNull(advisors, "advisors cannot be null");
|
||||
Assert.noNullElements(advisors, "advisors cannot contain null elements");
|
||||
this.advisors.addAll(advisors);
|
||||
return this;
|
||||
}
|
||||
@@ -334,57 +390,80 @@ public class DefaultChatClient implements ChatClient {
|
||||
private final DefaultChatClientRequestSpec request;
|
||||
|
||||
public DefaultCallResponseSpec(DefaultChatClientRequestSpec request) {
|
||||
Assert.notNull(request, "request cannot be null");
|
||||
this.request = request;
|
||||
}
|
||||
|
||||
@Override
|
||||
public <T> ResponseEntity<ChatResponse, T> responseEntity(Class<T> type) {
|
||||
Assert.notNull(type, "the class must be non-null");
|
||||
Assert.notNull(type, "type cannot be null");
|
||||
return doResponseEntity(new BeanOutputConverter<T>(type));
|
||||
}
|
||||
|
||||
@Override
|
||||
public <T> ResponseEntity<ChatResponse, T> responseEntity(ParameterizedTypeReference<T> type) {
|
||||
Assert.notNull(type, "type cannot be null");
|
||||
return doResponseEntity(new BeanOutputConverter<T>(type));
|
||||
}
|
||||
|
||||
@Override
|
||||
public <T> ResponseEntity<ChatResponse, T> responseEntity(
|
||||
StructuredOutputConverter<T> structuredOutputConverter) {
|
||||
Assert.notNull(structuredOutputConverter, "structuredOutputConverter cannot be null");
|
||||
return doResponseEntity(structuredOutputConverter);
|
||||
}
|
||||
|
||||
protected <T> ResponseEntity<ChatResponse, T> doResponseEntity(StructuredOutputConverter<T> boc) {
|
||||
var chatResponse = doGetObservableChatResponse(this.request, boc.getFormat());
|
||||
var responseContent = chatResponse.getResult().getOutput().getContent();
|
||||
T entity = boc.convert(responseContent);
|
||||
|
||||
protected <T> ResponseEntity<ChatResponse, T> doResponseEntity(StructuredOutputConverter<T> outputConverter) {
|
||||
Assert.notNull(outputConverter, "structuredOutputConverter cannot be null");
|
||||
var chatResponse = doGetObservableChatResponse(this.request, outputConverter.getFormat());
|
||||
var responseContent = getContentFromChatResponse(chatResponse);
|
||||
if (responseContent == null) {
|
||||
return new ResponseEntity<>(chatResponse, null);
|
||||
}
|
||||
T entity = outputConverter.convert(responseContent);
|
||||
return new ResponseEntity<>(chatResponse, entity);
|
||||
}
|
||||
|
||||
@Override
|
||||
@Nullable
|
||||
public <T> T entity(ParameterizedTypeReference<T> type) {
|
||||
return doSingleWithBeanOutputConverter(new BeanOutputConverter<T>(type));
|
||||
Assert.notNull(type, "type cannot be null");
|
||||
return doSingleWithBeanOutputConverter(new BeanOutputConverter<>(type));
|
||||
}
|
||||
|
||||
@Override
|
||||
@Nullable
|
||||
public <T> T entity(StructuredOutputConverter<T> structuredOutputConverter) {
|
||||
Assert.notNull(structuredOutputConverter, "structuredOutputConverter cannot be null");
|
||||
return doSingleWithBeanOutputConverter(structuredOutputConverter);
|
||||
}
|
||||
|
||||
private <T> T doSingleWithBeanOutputConverter(StructuredOutputConverter<T> boc) {
|
||||
var chatResponse = doGetObservableChatResponse(this.request, boc.getFormat());
|
||||
var stringResponse = chatResponse.getResult().getOutput().getContent();
|
||||
return boc.convert(stringResponse);
|
||||
}
|
||||
|
||||
@Override
|
||||
@Nullable
|
||||
public <T> T entity(Class<T> type) {
|
||||
Assert.notNull(type, "the class must be non-null");
|
||||
var boc = new BeanOutputConverter<T>(type);
|
||||
return doSingleWithBeanOutputConverter(boc);
|
||||
Assert.notNull(type, "type cannot be null");
|
||||
var outputConverter = new BeanOutputConverter<>(type);
|
||||
return doSingleWithBeanOutputConverter(outputConverter);
|
||||
}
|
||||
|
||||
@Nullable
|
||||
private <T> T doSingleWithBeanOutputConverter(StructuredOutputConverter<T> outputConverter) {
|
||||
var chatResponse = doGetObservableChatResponse(this.request, outputConverter.getFormat());
|
||||
var stringResponse = getContentFromChatResponse(chatResponse);
|
||||
if (stringResponse == null) {
|
||||
return null;
|
||||
}
|
||||
return outputConverter.convert(stringResponse);
|
||||
}
|
||||
|
||||
@Nullable
|
||||
private ChatResponse doGetChatResponse() {
|
||||
return this.doGetObservableChatResponse(this.request, "");
|
||||
return this.doGetObservableChatResponse(this.request, null);
|
||||
}
|
||||
|
||||
@Nullable
|
||||
private ChatResponse doGetObservableChatResponse(DefaultChatClientRequestSpec inputRequest,
|
||||
String formatParam) {
|
||||
@Nullable String formatParam) {
|
||||
|
||||
ChatClientObservationContext observationContext = ChatClientObservationContext.builder()
|
||||
.withRequest(inputRequest)
|
||||
@@ -395,19 +474,15 @@ public class DefaultChatClient implements ChatClient {
|
||||
var observation = ChatClientObservationDocumentation.AI_CHAT_CLIENT.observation(
|
||||
inputRequest.getCustomObservationConvention(), DEFAULT_CHAT_CLIENT_OBSERVATION_CONVENTION,
|
||||
() -> observationContext, inputRequest.getObservationRegistry());
|
||||
return observation.observe(() -> {
|
||||
ChatResponse chatResponse = doGetChatResponse(inputRequest, formatParam, observation);
|
||||
return chatResponse;
|
||||
});
|
||||
|
||||
return observation.observe(() -> doGetChatResponse(inputRequest, formatParam, observation));
|
||||
}
|
||||
|
||||
private ChatResponse doGetChatResponse(DefaultChatClientRequestSpec inputRequestSpec, String formatParam,
|
||||
Observation parentObservation) {
|
||||
private ChatResponse doGetChatResponse(DefaultChatClientRequestSpec inputRequestSpec,
|
||||
@Nullable String formatParam, Observation parentObservation) {
|
||||
|
||||
AdvisedRequest advisedRequest = toAdvisedRequest(inputRequestSpec, formatParam);
|
||||
|
||||
// Apply the around advisor chain that terminates with the, last, model call
|
||||
// Apply the around advisor chain that terminates with the last model call
|
||||
// advisor.
|
||||
AdvisedResponse advisedResponse = inputRequestSpec.aroundAdvisorChainBuilder.build()
|
||||
.nextAroundCall(advisedRequest);
|
||||
@@ -415,12 +490,26 @@ public class DefaultChatClient implements ChatClient {
|
||||
return advisedResponse.response();
|
||||
}
|
||||
|
||||
@Nullable
|
||||
private static String getContentFromChatResponse(@Nullable ChatResponse chatResponse) {
|
||||
if (chatResponse == null || chatResponse.getResult() == null || chatResponse.getResult().getOutput() == null
|
||||
|| chatResponse.getResult().getOutput().getContent() == null) {
|
||||
return null;
|
||||
}
|
||||
return chatResponse.getResult().getOutput().getContent();
|
||||
}
|
||||
|
||||
@Override
|
||||
@Nullable
|
||||
public ChatResponse chatResponse() {
|
||||
return doGetChatResponse();
|
||||
}
|
||||
|
||||
@Override
|
||||
@Nullable
|
||||
public String content() {
|
||||
return doGetChatResponse().getResult().getOutput().getContent();
|
||||
ChatResponse chatResponse = doGetChatResponse();
|
||||
return getContentFromChatResponse(chatResponse);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -430,6 +519,7 @@ public class DefaultChatClient implements ChatClient {
|
||||
private final DefaultChatClientRequestSpec request;
|
||||
|
||||
public DefaultStreamResponseSpec(DefaultChatClientRequestSpec request) {
|
||||
Assert.notNull(request, "request cannot be null");
|
||||
this.request = request;
|
||||
}
|
||||
|
||||
@@ -448,11 +538,10 @@ public class DefaultChatClient implements ChatClient {
|
||||
observation.parentObservation(contextView.getOrDefault(ObservationThreadLocalAccessor.KEY, null))
|
||||
.start();
|
||||
|
||||
var initialAdvisedRequest = toAdvisedRequest(inputRequest, "");
|
||||
var initialAdvisedRequest = toAdvisedRequest(inputRequest, null);
|
||||
|
||||
// @formatter:off
|
||||
// Apply the around advisor chain that terminates with the, last,
|
||||
// model call advisor.
|
||||
// Apply the around advisor chain that terminates with the last model call advisor.
|
||||
Flux<AdvisedResponse> stream = inputRequest.aroundAdvisorChainBuilder.build().nextAroundStream(initialAdvisedRequest);
|
||||
|
||||
return stream
|
||||
@@ -464,10 +553,12 @@ public class DefaultChatClient implements ChatClient {
|
||||
});
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatResponse> chatResponse() {
|
||||
return doGetObservableFluxChatResponse(this.request);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<String> content() {
|
||||
return doGetObservableFluxChatResponse(this.request).map(r -> {
|
||||
if (r.getResult() == null || r.getResult().getOutput() == null
|
||||
@@ -508,10 +599,13 @@ public class DefaultChatClient implements ChatClient {
|
||||
|
||||
private final Map<String, Object> toolContext = new HashMap<>();
|
||||
|
||||
private String userText = "";
|
||||
@Nullable
|
||||
private String userText;
|
||||
|
||||
private String systemText = "";
|
||||
@Nullable
|
||||
private String systemText;
|
||||
|
||||
@Nullable
|
||||
private ChatOptions chatOptions;
|
||||
|
||||
/* copy constructor */
|
||||
@@ -521,11 +615,25 @@ public class DefaultChatClient implements ChatClient {
|
||||
ccr.observationRegistry, ccr.customObservationConvention, ccr.toolContext);
|
||||
}
|
||||
|
||||
public DefaultChatClientRequestSpec(ChatModel chatModel, String userText, Map<String, Object> userParams,
|
||||
String systemText, Map<String, Object> systemParams, List<FunctionCallback> functionCallbacks,
|
||||
List<Message> messages, List<String> functionNames, List<Media> media, ChatOptions chatOptions,
|
||||
List<Advisor> advisors, Map<String, Object> advisorParams, ObservationRegistry observationRegistry,
|
||||
ChatClientObservationConvention customObservationConvention, Map<String, Object> toolContext) {
|
||||
public DefaultChatClientRequestSpec(ChatModel chatModel, @Nullable String userText,
|
||||
Map<String, Object> userParams, @Nullable String systemText, Map<String, Object> systemParams,
|
||||
List<FunctionCallback> functionCallbacks, List<Message> messages, List<String> functionNames,
|
||||
List<Media> media, @Nullable ChatOptions chatOptions, List<Advisor> advisors,
|
||||
Map<String, Object> advisorParams, ObservationRegistry observationRegistry,
|
||||
@Nullable ChatClientObservationConvention customObservationConvention,
|
||||
Map<String, Object> toolContext) {
|
||||
|
||||
Assert.notNull(chatModel, "chatModel cannot be null");
|
||||
Assert.notNull(userParams, "userParams cannot be null");
|
||||
Assert.notNull(systemParams, "systemParams cannot be null");
|
||||
Assert.notNull(functionCallbacks, "functionCallbacks cannot be null");
|
||||
Assert.notNull(messages, "messages cannot be null");
|
||||
Assert.notNull(functionNames, "functionNames cannot be null");
|
||||
Assert.notNull(media, "media cannot be null");
|
||||
Assert.notNull(advisors, "advisors cannot be null");
|
||||
Assert.notNull(advisorParams, "advisorParams cannot be null");
|
||||
Assert.notNull(observationRegistry, "observationRegistry cannot be null");
|
||||
Assert.notNull(toolContext, "toolContext cannot be null");
|
||||
|
||||
this.chatModel = chatModel;
|
||||
this.chatOptions = chatOptions != null ? chatOptions.copy()
|
||||
@@ -543,7 +651,8 @@ public class DefaultChatClient implements ChatClient {
|
||||
this.advisors.addAll(advisors);
|
||||
this.advisorParams.putAll(advisorParams);
|
||||
this.observationRegistry = observationRegistry;
|
||||
this.customObservationConvention = customObservationConvention;
|
||||
this.customObservationConvention = customObservationConvention != null ? customObservationConvention
|
||||
: DEFAULT_CHAT_CLIENT_OBSERVATION_CONVENTION;
|
||||
this.toolContext.putAll(toolContext);
|
||||
|
||||
// @formatter:off
|
||||
@@ -600,6 +709,7 @@ public class DefaultChatClient implements ChatClient {
|
||||
return this.customObservationConvention;
|
||||
}
|
||||
|
||||
@Nullable
|
||||
public String getUserText() {
|
||||
return this.userText;
|
||||
}
|
||||
@@ -608,6 +718,7 @@ public class DefaultChatClient implements ChatClient {
|
||||
return this.userParams;
|
||||
}
|
||||
|
||||
@Nullable
|
||||
public String getSystemText() {
|
||||
return this.systemText;
|
||||
}
|
||||
@@ -616,6 +727,7 @@ public class DefaultChatClient implements ChatClient {
|
||||
return this.systemParams;
|
||||
}
|
||||
|
||||
@Nullable
|
||||
public ChatOptions getChatOptions() {
|
||||
return this.chatOptions;
|
||||
}
|
||||
@@ -655,13 +767,21 @@ public class DefaultChatClient implements ChatClient {
|
||||
public Builder mutate() {
|
||||
DefaultChatClientBuilder builder = (DefaultChatClientBuilder) ChatClient
|
||||
.builder(this.chatModel, this.observationRegistry, this.customObservationConvention)
|
||||
.defaultSystem(s -> s.text(this.systemText).params(this.systemParams))
|
||||
.defaultUser(u -> u.text(this.userText)
|
||||
.params(this.userParams)
|
||||
.media(this.media.toArray(new Media[this.media.size()])))
|
||||
.defaultOptions(this.chatOptions)
|
||||
.defaultFunctions(StringUtils.toStringArray(this.functionNames));
|
||||
|
||||
if (StringUtils.hasText(this.userText)) {
|
||||
builder.defaultUser(
|
||||
u -> u.text(this.userText).params(this.userParams).media(this.media.toArray(new Media[0])));
|
||||
}
|
||||
|
||||
if (StringUtils.hasText(this.systemText)) {
|
||||
builder.defaultSystem(s -> s.text(this.systemText).params(this.systemParams));
|
||||
}
|
||||
|
||||
if (this.chatOptions != null) {
|
||||
builder.defaultOptions(this.chatOptions);
|
||||
}
|
||||
|
||||
// workaround to set the missing fields.
|
||||
builder.defaultRequest.getMessages().addAll(this.messages);
|
||||
builder.defaultRequest.getFunctionCallbacks().addAll(this.functionCallbacks);
|
||||
@@ -671,43 +791,47 @@ public class DefaultChatClient implements ChatClient {
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec advisors(Consumer<ChatClient.AdvisorSpec> consumer) {
|
||||
Assert.notNull(consumer, "the consumer must be non-null");
|
||||
var as = new DefaultAdvisorSpec();
|
||||
consumer.accept(as);
|
||||
this.advisorParams.putAll(as.getParams());
|
||||
this.advisors.addAll(as.getAdvisors());
|
||||
this.aroundAdvisorChainBuilder.pushAll(as.getAdvisors());
|
||||
Assert.notNull(consumer, "consumer cannot be null");
|
||||
var advisorSpec = new DefaultAdvisorSpec();
|
||||
consumer.accept(advisorSpec);
|
||||
this.advisorParams.putAll(advisorSpec.getParams());
|
||||
this.advisors.addAll(advisorSpec.getAdvisors());
|
||||
this.aroundAdvisorChainBuilder.pushAll(advisorSpec.getAdvisors());
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec advisors(Advisor... advisors) {
|
||||
Assert.notNull(advisors, "the advisors must be non-null");
|
||||
Assert.notNull(advisors, "advisors cannot be null");
|
||||
Assert.noNullElements(advisors, "advisors cannot contain null elements");
|
||||
this.advisors.addAll(Arrays.asList(advisors));
|
||||
this.aroundAdvisorChainBuilder.pushAll(Arrays.asList(advisors));
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec advisors(List<Advisor> advisors) {
|
||||
Assert.notNull(advisors, "the advisors must be non-null");
|
||||
Assert.notNull(advisors, "advisors cannot be null");
|
||||
Assert.noNullElements(advisors, "advisors cannot contain null elements");
|
||||
this.advisors.addAll(advisors);
|
||||
this.aroundAdvisorChainBuilder.pushAll(advisors);
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec messages(Message... messages) {
|
||||
Assert.notNull(messages, "the messages must be non-null");
|
||||
Assert.notNull(messages, "messages cannot be null");
|
||||
Assert.noNullElements(messages, "messages cannot contain null elements");
|
||||
this.messages.addAll(List.of(messages));
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec messages(List<Message> messages) {
|
||||
Assert.notNull(messages, "the messages must be non-null");
|
||||
Assert.notNull(messages, "messages cannot be null");
|
||||
Assert.noNullElements(messages, "messages cannot contain null elements");
|
||||
this.messages.addAll(messages);
|
||||
return this;
|
||||
}
|
||||
|
||||
public <T extends ChatOptions> ChatClientRequestSpec options(T options) {
|
||||
Assert.notNull(options, "the options must be non-null");
|
||||
Assert.notNull(options, "options cannot be null");
|
||||
this.chatOptions = options;
|
||||
return this;
|
||||
}
|
||||
@@ -720,9 +844,9 @@ public class DefaultChatClient implements ChatClient {
|
||||
public <I, O> ChatClientRequestSpec function(String name, String description,
|
||||
java.util.function.BiFunction<I, ToolContext, O> biFunction) {
|
||||
|
||||
Assert.hasText(name, "the name must be non-null and non-empty");
|
||||
Assert.hasText(description, "the description must be non-null and non-empty");
|
||||
Assert.notNull(biFunction, "the biFunction must be non-null");
|
||||
Assert.hasText(name, "name cannot be null or empty");
|
||||
Assert.hasText(description, "description cannot be null or empty");
|
||||
Assert.notNull(biFunction, "biFunction cannot be null");
|
||||
|
||||
FunctionCallbackWrapper<I, O> fcw = FunctionCallbackWrapper.builder(biFunction)
|
||||
.withDescription(description)
|
||||
@@ -733,12 +857,12 @@ public class DefaultChatClient implements ChatClient {
|
||||
return this;
|
||||
}
|
||||
|
||||
public <I, O> ChatClientRequestSpec function(String name, String description, Class<I> inputType,
|
||||
public <I, O> ChatClientRequestSpec function(String name, String description, @Nullable Class<I> inputType,
|
||||
java.util.function.Function<I, O> function) {
|
||||
|
||||
Assert.hasText(name, "the name must be non-null and non-empty");
|
||||
Assert.hasText(description, "the description must be non-null and non-empty");
|
||||
Assert.notNull(function, "the function must be non-null");
|
||||
Assert.hasText(name, "name cannot be null or empty");
|
||||
Assert.hasText(description, "description cannot be null or empty");
|
||||
Assert.notNull(function, "function cannot be null");
|
||||
|
||||
var fcw = FunctionCallbackWrapper.builder(function)
|
||||
.withDescription(description)
|
||||
@@ -751,36 +875,39 @@ public class DefaultChatClient implements ChatClient {
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec functions(String... functionBeanNames) {
|
||||
Assert.notNull(functionBeanNames, "the functionBeanNames must be non-null");
|
||||
Assert.notNull(functionBeanNames, "functionBeanNames cannot be null");
|
||||
Assert.noNullElements(functionBeanNames, "functionBeanNames cannot contain null elements");
|
||||
this.functionNames.addAll(List.of(functionBeanNames));
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec functions(FunctionCallback... functionCallbacks) {
|
||||
Assert.notNull(functionCallbacks, "the functionCallbacks must be non-null");
|
||||
Assert.notNull(functionCallbacks, "functionCallbacks cannot be null");
|
||||
Assert.noNullElements(functionCallbacks, "functionCallbacks cannot contain null elements");
|
||||
this.functionCallbacks.addAll(Arrays.asList(functionCallbacks));
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec toolContext(Map<String, Object> toolContext) {
|
||||
Assert.notNull(toolContext, "the toolContext must be non-null");
|
||||
Assert.notNull(toolContext, "toolContext cannot be null");
|
||||
Assert.noNullElements(toolContext.keySet(), "toolContext keys cannot contain null elements");
|
||||
Assert.noNullElements(toolContext.values(), "toolContext values cannot contain null elements");
|
||||
this.toolContext.putAll(toolContext);
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec system(String text) {
|
||||
Assert.notNull(text, "the text must be non-null");
|
||||
Assert.hasText(text, "text cannot be null or empty");
|
||||
this.systemText = text;
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec system(Resource textResource, Charset charset) {
|
||||
|
||||
Assert.notNull(textResource, "the text resource must be non-null");
|
||||
Assert.notNull(charset, "the charset must be non-null");
|
||||
public ChatClientRequestSpec system(Resource text, Charset charset) {
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
Assert.notNull(charset, "charset cannot be null");
|
||||
|
||||
try {
|
||||
this.systemText = textResource.getContentAsString(charset);
|
||||
this.systemText = text.getContentAsString(charset);
|
||||
}
|
||||
catch (IOException e) {
|
||||
throw new RuntimeException(e);
|
||||
@@ -789,32 +916,30 @@ public class DefaultChatClient implements ChatClient {
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec system(Resource text) {
|
||||
Assert.notNull(text, "the text resource must be non-null");
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
return this.system(text, Charset.defaultCharset());
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec system(Consumer<PromptSystemSpec> consumer) {
|
||||
Assert.notNull(consumer, "consumer cannot be null");
|
||||
|
||||
Assert.notNull(consumer, "the consumer must be non-null");
|
||||
|
||||
var ss = new DefaultPromptSystemSpec();
|
||||
consumer.accept(ss);
|
||||
this.systemText = StringUtils.hasText(ss.text()) ? ss.text() : this.systemText;
|
||||
this.systemParams.putAll(ss.params());
|
||||
var systemSpec = new DefaultPromptSystemSpec();
|
||||
consumer.accept(systemSpec);
|
||||
this.systemText = StringUtils.hasText(systemSpec.text()) ? systemSpec.text() : this.systemText;
|
||||
this.systemParams.putAll(systemSpec.params());
|
||||
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec user(String text) {
|
||||
Assert.notNull(text, "the text must be non-null");
|
||||
Assert.hasText(text, "text cannot be null or empty");
|
||||
this.userText = text;
|
||||
return this;
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec user(Resource text, Charset charset) {
|
||||
|
||||
Assert.notNull(text, "the text resource must be non-null");
|
||||
Assert.notNull(charset, "the charset must be non-null");
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
Assert.notNull(charset, "charset cannot be null");
|
||||
|
||||
try {
|
||||
this.userText = text.getContentAsString(charset);
|
||||
@@ -826,12 +951,12 @@ public class DefaultChatClient implements ChatClient {
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec user(Resource text) {
|
||||
Assert.notNull(text, "the text resource must be non-null");
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
return this.user(text, Charset.defaultCharset());
|
||||
}
|
||||
|
||||
public ChatClientRequestSpec user(Consumer<PromptUserSpec> consumer) {
|
||||
Assert.notNull(consumer, "the consumer must be non-null");
|
||||
Assert.notNull(consumer, "consumer cannot be null");
|
||||
|
||||
var us = new DefaultPromptUserSpec();
|
||||
consumer.accept(us);
|
||||
@@ -860,6 +985,8 @@ public class DefaultChatClient implements ChatClient {
|
||||
private final Prompt prompt;
|
||||
|
||||
public DefaultCallPromptResponseSpec(ChatModel chatModel, Prompt prompt) {
|
||||
Assert.notNull(chatModel, "chatModel cannot be null");
|
||||
Assert.notNull(prompt, "prompt cannot be null");
|
||||
this.chatModel = chatModel;
|
||||
this.prompt = prompt;
|
||||
}
|
||||
@@ -889,6 +1016,8 @@ public class DefaultChatClient implements ChatClient {
|
||||
private final StreamingChatModel chatModel;
|
||||
|
||||
public DefaultStreamPromptResponseSpec(StreamingChatModel streamingChatModel, Prompt prompt) {
|
||||
Assert.notNull(streamingChatModel, "streamingChatModel cannot be null");
|
||||
Assert.notNull(prompt, "prompt cannot be null");
|
||||
this.chatModel = streamingChatModel;
|
||||
this.prompt = prompt;
|
||||
}
|
||||
|
||||
@@ -35,6 +35,7 @@ import org.springframework.ai.chat.model.ToolContext;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.model.function.FunctionCallback;
|
||||
import org.springframework.core.io.Resource;
|
||||
import org.springframework.lang.Nullable;
|
||||
import org.springframework.util.Assert;
|
||||
|
||||
/**
|
||||
@@ -46,6 +47,7 @@ import org.springframework.util.Assert;
|
||||
* @author Christian Tzolov
|
||||
* @author Josh Long
|
||||
* @author Arjen Poutsma
|
||||
* @author Thomas Vitale
|
||||
* @since 1.0.0
|
||||
*/
|
||||
public class DefaultChatClientBuilder implements Builder {
|
||||
@@ -57,10 +59,10 @@ public class DefaultChatClientBuilder implements Builder {
|
||||
}
|
||||
|
||||
public DefaultChatClientBuilder(ChatModel chatModel, ObservationRegistry observationRegistry,
|
||||
ChatClientObservationConvention customObservationConvention) {
|
||||
@Nullable ChatClientObservationConvention customObservationConvention) {
|
||||
Assert.notNull(chatModel, "the " + ChatModel.class.getName() + " must be non-null");
|
||||
Assert.notNull(observationRegistry, "the " + ObservationRegistry.class.getName() + " must be non-null");
|
||||
this.defaultRequest = new DefaultChatClientRequestSpec(chatModel, "", Map.of(), "", Map.of(), List.of(),
|
||||
this.defaultRequest = new DefaultChatClientRequestSpec(chatModel, null, Map.of(), null, Map.of(), List.of(),
|
||||
List.of(), List.of(), List.of(), null, List.of(), Map.of(), observationRegistry,
|
||||
customObservationConvention, Map.of());
|
||||
}
|
||||
@@ -69,8 +71,8 @@ public class DefaultChatClientBuilder implements Builder {
|
||||
return new DefaultChatClient(this.defaultRequest);
|
||||
}
|
||||
|
||||
public Builder defaultAdvisors(Advisor... advisor) {
|
||||
this.defaultRequest.advisors(advisor);
|
||||
public Builder defaultAdvisors(Advisor... advisors) {
|
||||
this.defaultRequest.advisors(advisors);
|
||||
return this;
|
||||
}
|
||||
|
||||
@@ -95,6 +97,8 @@ public class DefaultChatClientBuilder implements Builder {
|
||||
}
|
||||
|
||||
public Builder defaultUser(Resource text, Charset charset) {
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
Assert.notNull(charset, "charset cannot be null");
|
||||
try {
|
||||
this.defaultRequest.user(text.getContentAsString(charset));
|
||||
}
|
||||
@@ -119,6 +123,8 @@ public class DefaultChatClientBuilder implements Builder {
|
||||
}
|
||||
|
||||
public Builder defaultSystem(Resource text, Charset charset) {
|
||||
Assert.notNull(text, "text cannot be null");
|
||||
Assert.notNull(charset, "charset cannot be null");
|
||||
try {
|
||||
this.defaultRequest.system(text.getContentAsString(charset));
|
||||
}
|
||||
|
||||
@@ -16,6 +16,8 @@
|
||||
|
||||
package org.springframework.ai.chat.client;
|
||||
|
||||
import org.springframework.lang.Nullable;
|
||||
|
||||
/**
|
||||
* Represents a {@link org.springframework.ai.model.Model} response that includes the
|
||||
* entire response along withe specified response entity type.
|
||||
@@ -25,14 +27,17 @@ package org.springframework.ai.chat.client;
|
||||
* @param response the entire response object.
|
||||
* @param entity the converted entity object.
|
||||
* @author Christian Tzolov
|
||||
* @author Thomas Vitale
|
||||
* @since 1.0.0
|
||||
*/
|
||||
public record ResponseEntity<R, E>(R response, E entity) {
|
||||
public record ResponseEntity<R, E>(@Nullable R response, @Nullable E entity) {
|
||||
|
||||
@Nullable
|
||||
public R getResponse() {
|
||||
return this.response;
|
||||
}
|
||||
|
||||
@Nullable
|
||||
public E getEntity() {
|
||||
return this.entity;
|
||||
}
|
||||
|
||||
@@ -86,15 +86,30 @@ public record AdvisedRequest(
|
||||
Assert.notNull(chatModel, "chatModel cannot be null");
|
||||
Assert.hasText(userText, "userText cannot be null or empty");
|
||||
Assert.notNull(media, "media cannot be null");
|
||||
Assert.noNullElements(media, "media cannot contain null elements");
|
||||
Assert.notNull(functionNames, "functionNames cannot be null");
|
||||
Assert.noNullElements(functionNames, "functionNames cannot contain null elements");
|
||||
Assert.notNull(functionCallbacks, "functionCallbacks cannot be null");
|
||||
Assert.noNullElements(functionCallbacks, "functionCallbacks cannot contain null elements");
|
||||
Assert.notNull(messages, "messages cannot be null");
|
||||
Assert.noNullElements(messages, "messages cannot contain null elements");
|
||||
Assert.notNull(userParams, "userParams cannot be null");
|
||||
Assert.noNullElements(userParams.keySet(), "userParams keys cannot contain null elements");
|
||||
Assert.noNullElements(userParams.values(), "userParams values cannot contain null elements");
|
||||
Assert.notNull(systemParams, "systemParams cannot be null");
|
||||
Assert.noNullElements(systemParams.keySet(), "systemParams keys cannot contain null elements");
|
||||
Assert.noNullElements(systemParams.values(), "systemParams values cannot contain null elements");
|
||||
Assert.notNull(advisors, "advisors cannot be null");
|
||||
Assert.noNullElements(advisors, "advisors cannot contain null elements");
|
||||
Assert.notNull(advisorParams, "advisorParams cannot be null");
|
||||
Assert.noNullElements(advisorParams.keySet(), "advisorParams keys cannot contain null elements");
|
||||
Assert.noNullElements(advisorParams.values(), "advisorParams values cannot contain null elements");
|
||||
Assert.notNull(adviseContext, "adviseContext cannot be null");
|
||||
Assert.noNullElements(adviseContext.keySet(), "adviseContext keys cannot contain null elements");
|
||||
Assert.noNullElements(adviseContext.values(), "adviseContext values cannot contain null elements");
|
||||
Assert.notNull(toolContext, "toolContext cannot be null");
|
||||
Assert.noNullElements(toolContext.keySet(), "toolContext keys cannot contain null elements");
|
||||
Assert.noNullElements(toolContext.values(), "toolContext values cannot contain null elements");
|
||||
}
|
||||
|
||||
public static Builder builder() {
|
||||
|
||||
@@ -22,6 +22,7 @@ import java.util.Map;
|
||||
import java.util.function.Function;
|
||||
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.lang.Nullable;
|
||||
import org.springframework.util.Assert;
|
||||
|
||||
/**
|
||||
@@ -31,11 +32,12 @@ import org.springframework.util.Assert;
|
||||
* @author Thomas Vitale
|
||||
* @since 1.0.0
|
||||
*/
|
||||
public record AdvisedResponse(ChatResponse response, Map<String, Object> adviseContext) {
|
||||
public record AdvisedResponse(@Nullable ChatResponse response, Map<String, Object> adviseContext) {
|
||||
|
||||
public AdvisedResponse {
|
||||
Assert.notNull(response, "response cannot be null");
|
||||
Assert.notNull(adviseContext, "adviseContext cannot be null");
|
||||
Assert.noNullElements(adviseContext.keySet(), "adviseContext keys cannot be null");
|
||||
Assert.noNullElements(adviseContext.values(), "adviseContext values cannot be null");
|
||||
}
|
||||
|
||||
public static Builder builder() {
|
||||
@@ -55,6 +57,7 @@ public record AdvisedResponse(ChatResponse response, Map<String, Object> adviseC
|
||||
|
||||
public static final class Builder {
|
||||
|
||||
@Nullable
|
||||
private ChatResponse response;
|
||||
|
||||
private Map<String, Object> adviseContext;
|
||||
@@ -62,7 +65,7 @@ public record AdvisedResponse(ChatResponse response, Map<String, Object> adviseC
|
||||
private Builder() {
|
||||
}
|
||||
|
||||
public Builder withResponse(ChatResponse response) {
|
||||
public Builder withResponse(@Nullable ChatResponse response) {
|
||||
this.response = response;
|
||||
return this;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
/*
|
||||
* Copyright 2023-2024 the original author or authors.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* https://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
@NonNullApi
|
||||
@NonNullFields
|
||||
package org.springframework.ai.chat.client;
|
||||
|
||||
import org.springframework.lang.NonNullApi;
|
||||
import org.springframework.lang.NonNullFields;
|
||||
@@ -17,6 +17,7 @@
|
||||
package org.springframework.ai.chat.prompt;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Arrays;
|
||||
import java.util.Collections;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
@@ -54,6 +55,10 @@ public class Prompt implements ModelRequest<List<Message>> {
|
||||
this(messages, null);
|
||||
}
|
||||
|
||||
public Prompt(Message... messages) {
|
||||
this(Arrays.asList(messages), null);
|
||||
}
|
||||
|
||||
public Prompt(String contents, ChatOptions chatOptions) {
|
||||
this(new UserMessage(contents), chatOptions);
|
||||
}
|
||||
|
||||
@@ -33,6 +33,7 @@ import reactor.core.publisher.Flux;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.messages.MessageType;
|
||||
import org.springframework.ai.chat.messages.SystemMessage;
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
@@ -46,6 +47,7 @@ import org.springframework.core.io.DefaultResourceLoader;
|
||||
import org.springframework.util.MimeTypeUtils;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
import static org.assertj.core.api.Assertions.assertThatThrownBy;
|
||||
import static org.mockito.BDDMockito.given;
|
||||
|
||||
/**
|
||||
@@ -69,7 +71,7 @@ public class ChatClientTest {
|
||||
|
||||
// ChatClient Builder Tests
|
||||
@Test
|
||||
public void defaultSystemText() {
|
||||
void defaultSystemText() {
|
||||
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
@@ -118,7 +120,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void defaultSystemTextLambda() {
|
||||
void defaultSystemTextLambda() {
|
||||
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
@@ -194,7 +196,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void mutateDefaults() {
|
||||
void mutateDefaults() {
|
||||
|
||||
PortableFunctionCallingOptions options = new FunctionCallingOptionsBuilder().build();
|
||||
given(this.chatModel.getDefaultOptions()).willReturn(options);
|
||||
@@ -322,7 +324,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void mutatePrompt() {
|
||||
void mutatePrompt() {
|
||||
|
||||
PortableFunctionCallingOptions options = new FunctionCallingOptionsBuilder().build();
|
||||
given(this.chatModel.getDefaultOptions()).willReturn(options);
|
||||
@@ -412,7 +414,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void defaultUserText() {
|
||||
void defaultUserText() {
|
||||
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
@@ -437,7 +439,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void simpleUserPromptAsString() {
|
||||
void simpleUserPromptAsString() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
@@ -450,7 +452,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void simpleUserPrompt() {
|
||||
void simpleUserPrompt() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
@@ -463,7 +465,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void simpleUserPromptObject() throws MalformedURLException {
|
||||
void simpleUserPromptObject() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
@@ -482,7 +484,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void simpleSystemPrompt() throws MalformedURLException {
|
||||
void simpleSystemPrompt() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
@@ -503,7 +505,7 @@ public class ChatClientTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
public void complexCall() throws MalformedURLException {
|
||||
void complexCall() throws MalformedURLException {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
@@ -545,4 +547,264 @@ public class ChatClientTest {
|
||||
assertThat(options.getFunctions()).isEmpty();
|
||||
}
|
||||
|
||||
// Constructors
|
||||
|
||||
@Test
|
||||
void whenCreateAndChatModelIsNullThenThrow() {
|
||||
assertThatThrownBy(() -> ChatClient.create(null)).isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("chatModel cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenCreateAndObservationRegistryIsNullThenThrow() {
|
||||
assertThatThrownBy(() -> ChatClient.create(this.chatModel, null, null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("observationRegistry cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenBuilderAndChatModelIsNullThenThrow() {
|
||||
assertThatThrownBy(() -> ChatClient.builder(null)).isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("chatModel cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenBuilderAndObservationRegistryIsNullThenThrow() {
|
||||
assertThatThrownBy(() -> ChatClient.builder(this.chatModel, null, null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("observationRegistry cannot be null");
|
||||
}
|
||||
|
||||
// Prompt Tests - User
|
||||
|
||||
@Test
|
||||
void whenPromptWithStringContent() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
var content = chatClient.prompt("my question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(1);
|
||||
var userMessage = this.promptCaptor.getValue().getInstructions().get(0);
|
||||
assertThat(userMessage.getContent()).isEqualTo("my question");
|
||||
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenPromptWithMessages() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
var prompt = new Prompt(new SystemMessage("instructions"), new UserMessage("my question"));
|
||||
var content = chatClient.prompt(prompt).call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(2);
|
||||
var userMessage = this.promptCaptor.getValue().getInstructions().get(1);
|
||||
assertThat(userMessage.getContent()).isEqualTo("my question");
|
||||
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenPromptWithStringContentAndUserText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
var content = chatClient.prompt("my question").user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(2);
|
||||
var userMessage = this.promptCaptor.getValue().getInstructions().get(1);
|
||||
assertThat(userMessage.getContent()).isEqualTo("another question");
|
||||
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenPromptWithHistoryAndUserText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
var prompt = new Prompt(new UserMessage("my question"), new AssistantMessage("your answer"));
|
||||
var content = chatClient.prompt(prompt).user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(3);
|
||||
var userMessage = this.promptCaptor.getValue().getInstructions().get(2);
|
||||
assertThat(userMessage.getContent()).isEqualTo("another question");
|
||||
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenPromptWithUserMessageAndUserText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
var prompt = new Prompt(new UserMessage("my question"));
|
||||
var content = chatClient.prompt(prompt).user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(2);
|
||||
var userMessage = this.promptCaptor.getValue().getInstructions().get(1);
|
||||
assertThat(userMessage.getContent()).isEqualTo("another question");
|
||||
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenMessagesWithHistoryAndUserText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
List<Message> messages = List.of(new UserMessage("my question"), new AssistantMessage("your answer"));
|
||||
var content = chatClient.prompt().messages(messages).user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(3);
|
||||
var userMessage = this.promptCaptor.getValue().getInstructions().get(2);
|
||||
assertThat(userMessage.getContent()).isEqualTo("another question");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenMessagesWithUserMessageAndUserText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
List<Message> messages = List.of(new UserMessage("my question"));
|
||||
var content = chatClient.prompt().messages(messages).user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(2);
|
||||
var userMessage = this.promptCaptor.getValue().getInstructions().get(1);
|
||||
assertThat(userMessage.getContent()).isEqualTo("another question");
|
||||
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
|
||||
}
|
||||
|
||||
// Prompt Tests - System
|
||||
|
||||
@Test
|
||||
void whenPromptWithMessagesAndSystemText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
var prompt = new Prompt(new UserMessage("my question"), new AssistantMessage("your answer"));
|
||||
var content = chatClient.prompt(prompt).system("instructions").user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(4);
|
||||
var systemMessage = this.promptCaptor.getValue().getInstructions().get(2);
|
||||
assertThat(systemMessage.getContent()).isEqualTo("instructions");
|
||||
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenPromptWithSystemMessageAndNoSystemText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
var prompt = new Prompt(new SystemMessage("instructions"), new UserMessage("my question"));
|
||||
var content = chatClient.prompt(prompt).user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(3);
|
||||
var systemMessage = this.promptCaptor.getValue().getInstructions().get(0);
|
||||
assertThat(systemMessage.getContent()).isEqualTo("instructions");
|
||||
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenPromptWithSystemMessageAndSystemText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
var prompt = new Prompt(new SystemMessage("instructions"), new UserMessage("my question"));
|
||||
var content = chatClient.prompt(prompt).system("other instructions").user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(4);
|
||||
var systemMessage = this.promptCaptor.getValue().getInstructions().get(2);
|
||||
assertThat(systemMessage.getContent()).isEqualTo("other instructions");
|
||||
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenMessagesAndSystemText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
List<Message> messages = List.of(new UserMessage("my question"), new AssistantMessage("your answer"));
|
||||
var content = chatClient.prompt()
|
||||
.messages(messages)
|
||||
.system("instructions")
|
||||
.user("another question")
|
||||
.call()
|
||||
.content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(4);
|
||||
var systemMessage = this.promptCaptor.getValue().getInstructions().get(2);
|
||||
assertThat(systemMessage.getContent()).isEqualTo("instructions");
|
||||
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenMessagesWithSystemMessageAndNoSystemText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
List<Message> messages = List.of(new SystemMessage("instructions"), new UserMessage("my question"));
|
||||
var content = chatClient.prompt().messages(messages).user("another question").call().content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(3);
|
||||
var systemMessage = this.promptCaptor.getValue().getInstructions().get(0);
|
||||
assertThat(systemMessage.getContent()).isEqualTo("instructions");
|
||||
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenMessagesWithSystemMessageAndSystemText() {
|
||||
given(this.chatModel.call(this.promptCaptor.capture()))
|
||||
.willReturn(new ChatResponse(List.of(new Generation(new AssistantMessage("response")))));
|
||||
|
||||
var chatClient = ChatClient.builder(this.chatModel).build();
|
||||
List<Message> messages = List.of(new SystemMessage("instructions"), new UserMessage("my question"));
|
||||
var content = chatClient.prompt()
|
||||
.messages(messages)
|
||||
.system("other instructions")
|
||||
.user("another question")
|
||||
.call()
|
||||
.content();
|
||||
|
||||
assertThat(content).isEqualTo("response");
|
||||
|
||||
assertThat(this.promptCaptor.getValue().getInstructions()).hasSize(4);
|
||||
var systemMessage = this.promptCaptor.getValue().getInstructions().get(2);
|
||||
assertThat(systemMessage.getContent()).isEqualTo("other instructions");
|
||||
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
/*
|
||||
* Copyright 2023-2024 the original author or authors.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* https://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
package org.springframework.ai.chat.client;
|
||||
|
||||
import java.nio.charset.Charset;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.core.io.ClassPathResource;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThatThrownBy;
|
||||
import static org.mockito.Mockito.mock;
|
||||
|
||||
/**
|
||||
* Unit tests for {@link DefaultChatClientBuilder}.
|
||||
*
|
||||
* @author Thomas Vitale
|
||||
*/
|
||||
class DefaultChatClientBuilderTests {
|
||||
|
||||
@Test
|
||||
void whenChatModelIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new DefaultChatClientBuilder(null)).isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("the org.springframework.ai.chat.model.ChatModel must be non-null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenObservationRegistryIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new DefaultChatClientBuilder(mock(ChatModel.class), null, null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("the io.micrometer.observation.ObservationRegistry must be non-null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenUserResourceIsNullThenThrows() {
|
||||
DefaultChatClientBuilder builder = new DefaultChatClientBuilder(mock(ChatModel.class));
|
||||
assertThatThrownBy(() -> builder.defaultUser(null, Charset.defaultCharset()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("text cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenUserCharsetIsNullThenThrows() {
|
||||
DefaultChatClientBuilder builder = new DefaultChatClientBuilder(mock(ChatModel.class));
|
||||
assertThatThrownBy(() -> builder.defaultUser(new ClassPathResource("user-prompt.txt"), null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("charset cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenSystemResourceIsNullThenThrows() {
|
||||
DefaultChatClientBuilder builder = new DefaultChatClientBuilder(mock(ChatModel.class));
|
||||
assertThatThrownBy(() -> builder.defaultSystem(null, Charset.defaultCharset()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("text cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenSystemCharsetIsNullThenThrows() {
|
||||
DefaultChatClientBuilder builder = new DefaultChatClientBuilder(mock(ChatModel.class));
|
||||
assertThatThrownBy(() -> builder.defaultSystem(new ClassPathResource("system-prompt.txt"), null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("charset cannot be null");
|
||||
}
|
||||
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,148 @@
|
||||
/*
|
||||
* Copyright 2023-2024 the original author or authors.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* https://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
package org.springframework.ai.chat.client.advisor.api;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
import static org.assertj.core.api.Assertions.assertThatThrownBy;
|
||||
import static org.mockito.Mockito.mock;
|
||||
|
||||
/**
|
||||
* Unit tests for {@link AdvisedRequest}.
|
||||
*
|
||||
* @author Thomas Vitale
|
||||
*/
|
||||
class AdvisedRequestTests {
|
||||
|
||||
@Test
|
||||
void buildAdvisedRequest() {
|
||||
AdvisedRequest request = new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), Map.of(), Map.of(), List.of(), Map.of(), Map.of(), Map.of());
|
||||
assertThat(request).isNotNull();
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenChatModelIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(null, "user", null, null, List.of(), List.of(), List.of(),
|
||||
List.of(), Map.of(), Map.of(), List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("chatModel cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenUserTextIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), null, null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), Map.of(), Map.of(), List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("userText cannot be null or empty");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenUserTextIsEmptyThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "", null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), Map.of(), Map.of(), List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("userText cannot be null or empty");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenMediaIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, null, List.of(),
|
||||
List.of(), List.of(), Map.of(), Map.of(), List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("media cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenFunctionNamesIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), null,
|
||||
List.of(), List.of(), Map.of(), Map.of(), List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("functionNames cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenFunctionCallbacksIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
null, List.of(), Map.of(), Map.of(), List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("functionCallbacks cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenMessagesIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
List.of(), null, Map.of(), Map.of(), List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("messages cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenUserParamsIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), null, Map.of(), List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("userParams cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenSystemParamsIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), Map.of(), null, List.of(), Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("systemParams cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenAdvisorsIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), Map.of(), Map.of(), null, Map.of(), Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("advisors cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenAdvisorParamsIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), Map.of(), Map.of(), List.of(), null, Map.of(), Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("advisorParams cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenAdviseContextIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), Map.of(), Map.of(), List.of(), Map.of(), null, Map.of()))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("adviseContext cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenToolContextIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedRequest(mock(ChatModel.class), "user", null, null, List.of(), List.of(),
|
||||
List.of(), List.of(), Map.of(), Map.of(), List.of(), Map.of(), Map.of(), null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("toolContext cannot be null");
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
/*
|
||||
* Copyright 2023-2024 the original author or authors.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* https://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
package org.springframework.ai.chat.client.advisor.api;
|
||||
|
||||
import java.util.HashMap;
|
||||
import java.util.Map;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
import static org.assertj.core.api.Assertions.assertThatThrownBy;
|
||||
import static org.mockito.Mockito.mock;
|
||||
|
||||
/**
|
||||
* Unit tests for {@link AdvisedResponse}.
|
||||
*
|
||||
* @author Thomas Vitale
|
||||
*/
|
||||
class AdvisedResponseTests {
|
||||
|
||||
@Test
|
||||
void buildAdvisedResponse() {
|
||||
AdvisedResponse advisedResponse = new AdvisedResponse(mock(ChatResponse.class), Map.of());
|
||||
assertThat(advisedResponse).isNotNull();
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenAdviseContextIsNullThenThrows() {
|
||||
assertThatThrownBy(() -> new AdvisedResponse(mock(ChatResponse.class), null))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("adviseContext cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenAdviseContextKeysIsNullThenThrows() {
|
||||
Map<String, Object> adviseContext = new HashMap<>();
|
||||
adviseContext.put(null, "value");
|
||||
assertThatThrownBy(() -> new AdvisedResponse(mock(ChatResponse.class), adviseContext))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("adviseContext keys cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenAdviseContextValuesIsNullThenThrows() {
|
||||
Map<String, Object> adviseContext = new HashMap<>();
|
||||
adviseContext.put("key", null);
|
||||
assertThatThrownBy(() -> new AdvisedResponse(mock(ChatResponse.class), adviseContext))
|
||||
.isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("adviseContext values cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenBuildFromNullAdvisedResponseThenThrows() {
|
||||
assertThatThrownBy(() -> AdvisedResponse.from(null)).isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("advisedResponse cannot be null");
|
||||
}
|
||||
|
||||
@Test
|
||||
void buildFromAdvisedResponse() {
|
||||
AdvisedResponse advisedResponse = new AdvisedResponse(mock(ChatResponse.class), Map.of());
|
||||
AdvisedResponse.Builder builder = AdvisedResponse.from(advisedResponse);
|
||||
assertThat(builder).isNotNull();
|
||||
}
|
||||
|
||||
@Test
|
||||
void whenUpdateFromNullContextThenThrows() {
|
||||
AdvisedResponse advisedResponse = new AdvisedResponse(mock(ChatResponse.class), Map.of());
|
||||
assertThatThrownBy(() -> advisedResponse.updateContext(null)).isInstanceOf(IllegalArgumentException.class)
|
||||
.hasMessage("contextTransform cannot be null");
|
||||
}
|
||||
|
||||
}
|
||||
1
spring-ai-core/src/test/resources/system-prompt.txt
Normal file
1
spring-ai-core/src/test/resources/system-prompt.txt
Normal file
@@ -0,0 +1 @@
|
||||
instructions
|
||||
BIN
spring-ai-core/src/test/resources/tabby-cat.png
Normal file
BIN
spring-ai-core/src/test/resources/tabby-cat.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 240 KiB |
1
spring-ai-core/src/test/resources/user-prompt.txt
Normal file
1
spring-ai-core/src/test/resources/user-prompt.txt
Normal file
@@ -0,0 +1 @@
|
||||
my question
|
||||
Reference in New Issue
Block a user