Fix ChatClient list() convertions

This commit is contained in:
Christian Tzolov
2024-05-21 23:13:33 +02:00
parent bc5c47b201
commit 545bc892b8
5 changed files with 85 additions and 224 deletions

View File

@@ -34,7 +34,6 @@ import reactor.core.publisher.Flux;
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.converter.BeanOutputConverter;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.openai.OpenAiChatOptions;
import org.springframework.ai.openai.OpenAiTestConfiguration;
import org.springframework.ai.openai.api.OpenAiApi;
@@ -72,49 +71,51 @@ class OpenAiChatClientIT extends AbstractIT {
// @formatter:on
logger.info("" + response);
// UserMessage userMessage = new UserMessage(
// "Tell me about 3 famous pirates from the Golden Age of Piracy and what they
// did.");
// SystemPromptTemplate systemPromptTemplate = new
// SystemPromptTemplate(systemResource);
// Message systemMessage = systemPromptTemplate.createMessage(Map.of("name",
// "Bob", "voice", "pirate"));
// Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
// ChatResponse response = modelCaller.call(prompt);
assertThat(response.getResults()).hasSize(1);
assertThat(response.getResults().get(0).getOutput().getContent()).contains("Blackbeard");
// needs fine tuning... evaluateQuestionAndAnswer(request, response, false);
}
// @Test
@Test
void listOutputConverter() {
// TODO: there is a problem here.
// @formatter:off
Collection<String> list = ChatClient.builder(modelCaller).build().prompt()
Collection<String> collection = ChatClient.builder(modelCaller).build().prompt()
.user(u -> u.text("List five {subject}")
.param("subject", "ice cream flavors"))
.call()
.list(String.class);
// @formatter:on
// DefaultConversionService conversionService = new DefaultConversionService();
// ListOutputConverter outputConverter = new
// ListOutputConverter(conversionService);
assertThat(collection).hasSize(5);
}
// String format = outputConverter.getFormat();
// String template = """
// List five {subject}
// {format}
// """;
// PromptTemplate promptTemplate = new PromptTemplate(template,
// Map.of("subject", "ice cream flavors", "format", format));
// Prompt prompt = new Prompt(promptTemplate.createMessage());
// Generation generation = this.modelCaller.call(prompt).getResult();
@Test
void listOutputConverter2() {
// List<String> list =
// outputConverter.convert(generation.getOutput().getContent());
assertThat(list).hasSize(5);
// @formatter:off
List<ActorsFilmsRecord> actorsFilms = ChatClient.builder(modelCaller).build().prompt()
.user("Generate the filmography of 5 movies for Tom Hanks and Bill Murray.")
.call()
.single(new ParameterizedTypeReference<List<ActorsFilmsRecord>>() {
});
// @formatter:on
logger.info("" + actorsFilms);
assertThat(actorsFilms).hasSize(2);
}
@Test
void listOutputConverter3() {
// @formatter:off
Collection<ActorsFilmsRecord> actorsFilms = ChatClient.builder(modelCaller).build().prompt()
.user("Generate the filmography of 5 movies for Tom Hanks and Bill Murray.")
.call()
.list(ActorsFilmsRecord.class);
// @formatter:on
logger.info("" + actorsFilms);
assertThat(actorsFilms).hasSize(2);
}
@@ -124,25 +125,11 @@ class OpenAiChatClientIT extends AbstractIT {
Map<String, Object> result = ChatClient.builder(modelCaller).build().prompt()
.user(u -> u.text("Provide me a List of {subject}")
.param("subject", "an array of numbers from 1 to 9 under they key name 'numbers'"))
.call().single(new ParameterizedTypeReference<Map<String, Object>>() {
.call()
.single(new ParameterizedTypeReference<Map<String, Object>>() {
});
// @formatter:on
// MapOutputConverter outputConverter = new MapOutputConverter();
// String format = outputConverter.getFormat();
// String template = """
// Provide me a List of {subject}
// {format}
// """;
// PromptTemplate promptTemplate = new PromptTemplate(template,
// Map.of("subject", "an array of numbers from 1 to 9 under they key name
// 'numbers'", "format", format));
// Prompt prompt = new Prompt(promptTemplate.createMessage());
// Generation generation = modelCaller.call(prompt).getResult();
// Map<String, Object> result =
// outputConverter.convert(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -156,21 +143,6 @@ class OpenAiChatClientIT extends AbstractIT {
.single(ActorsFilms.class);
// @formatter:on
// BeanOutputConverter<ActorsFilms> outputConverter = new
// BeanOutputConverter<>(ActorsFilms.class);
// String format = outputConverter.getFormat();
// String template = """
// Generate the filmography for a random actor.
// {format}
// """;
// PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format",
// format));
// Prompt prompt = new Prompt(promptTemplate.createMessage());
// Generation generation = modelCaller.call(prompt).getResult();
// ActorsFilms actorsFilms =
// outputConverter.convert(generation.getOutput().getContent());
logger.info("" + actorsFilms);
assertThat(actorsFilms.getActor()).isNotBlank();
}
@@ -188,21 +160,6 @@ class OpenAiChatClientIT extends AbstractIT {
.single(ActorsFilmsRecord.class);
// @formatter:on
// BeanOutputConverter<ActorsFilmsRecord> outputConverter = new
// BeanOutputConverter<>(ActorsFilmsRecord.class);
// String format = outputConverter.getFormat();
// String template = """
// Generate the filmography of 5 movies for Tom Hanks.
// {format}
// """;
// PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format",
// format));
// Prompt prompt = new Prompt(promptTemplate.createMessage());
// Generation generation = modelCaller.call(prompt).getResult();
// ActorsFilmsRecord actorsFilms =
// outputConverter.convert(generation.getOutput().getContent());
logger.info("" + actorsFilms);
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
@@ -230,34 +187,8 @@ class OpenAiChatClientIT extends AbstractIT {
.collect(Collectors.joining());
// @formatter:on
// String generationTextFromStream = chatResponse.collectList()
// .block()
// .stream()
// .collect(Collectors.joining());
// BeanOutputConverter<ActorsFilmsRecord> outputConverter = new
// BeanOutputConverter<>(ActorsFilmsRecord.class);
// String format = outputConverter.getFormat();
// String template = """
// Generate the filmography of 5 movies for Tom Hanks.
// {format}
// """;
// PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format",
// format));
// Prompt prompt = new Prompt(promptTemplate.createMessage());
// String generationTextFromStream = streamingChatClient.stream(prompt)
// .collectList()
// .block()
// .stream()
// .map(ChatResponse::getResults)
// .flatMap(List::stream)
// .map(Generation::getOutput)
// .map(AssistantMessage::getContent)
// .collect(Collectors.joining());
ActorsFilmsRecord actorsFilms = outputConverter.convert(generationTextFromStream);
logger.info("" + actorsFilms);
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
@@ -274,23 +205,6 @@ class OpenAiChatClientIT extends AbstractIT {
.content();
// @formatter:on
// UserMessage userMessage = new UserMessage("What's the weather like in San
// Francisco, Tokyo, and Paris?");
// List<Message> messages = new ArrayList<>(List.of(userMessage));
// var promptOptions = OpenAiChatOptions.builder()
// .withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue())
// .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new
// MockWeatherService())
// .withName("getCurrentWeather")
// .withDescription("Get the weather in location")
// .withResponseConverter((response) -> "" + response.temp() + response.unit())
// .build()))
// .build();
// ChatResponse response = modelCaller.call(new Prompt(messages, promptOptions));
logger.info("Response: {}", response);
assertThat(response).containsAnyOf("30.0", "30");
@@ -309,24 +223,6 @@ class OpenAiChatClientIT extends AbstractIT {
.content();
// @formatter:on
// UserMessage userMessage = new UserMessage("What's the weather like in San
// Francisco, Tokyo, and Paris?");
// List<Message> messages = new ArrayList<>(List.of(userMessage));
// var promptOptions = OpenAiChatOptions.builder()
// // .withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue())
// .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new
// MockWeatherService())
// .withName("getCurrentWeather")
// .withDescription("Get the weather in location")
// .withResponseConverter((response) -> "" + response.temp() + response.unit())
// .build()))
// .build();
// Flux<ChatResponse> response = streamingChatClient.stream(new Prompt(messages,
// promptOptions));
String content = response.collectList().block().stream().collect(Collectors.joining());
logger.info("Response: {}", content);
@@ -350,15 +246,6 @@ class OpenAiChatClientIT extends AbstractIT {
.content();
// @formatter:on
// var imageData = new ClassPathResource("/test.png");
// var userMessage = new UserMessage("Explain what do you see on this picture?",
// List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData)));
// var response = modelCaller
// .call(new Prompt(List.of(userMessage),
// OpenAiChatOptions.builder().withModel(modelName).build()));
logger.info(response);
assertThat(response).contains("bananas", "apple");
assertThat(response).containsAnyOf("bowl", "basket");
@@ -372,9 +259,7 @@ class OpenAiChatClientIT extends AbstractIT {
URL url = new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png");
// @formatter:off
String response = ChatClient.builder(modelCaller)
.build()
.prompt()
String response = ChatClient.builder(modelCaller).build().prompt()
// TODO consider adding model(...) method to ChatClient as a shortcut to
// OpenAiChatOptions.builder().withModel(modelName).build()
.options(OpenAiChatOptions.builder().withModel(modelName).build())
@@ -383,16 +268,6 @@ class OpenAiChatClientIT extends AbstractIT {
.content();
// @formatter:on
// var userMessage = new UserMessage("Explain what do you see on this picture?",
// List
// .of(new Media(MimeTypeUtils.IMAGE_PNG,
// new
// URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"))));
// ChatResponse response = modelCaller
// .call(new Prompt(List.of(userMessage),
// OpenAiChatOptions.builder().withModel(modelName).build()));
logger.info(response);
assertThat(response).contains("bananas", "apple");
assertThat(response).containsAnyOf("bowl", "basket");
@@ -414,15 +289,6 @@ class OpenAiChatClientIT extends AbstractIT {
.content();
// @formatter:on
// var userMessage = new UserMessage("Explain what do you see on this picture?",
// List.of(new Media(MimeTypeUtils.IMAGE_PNG,
// new
// URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"))));
// Flux<ChatResponse> response = streamingChatClient.stream(new
// Prompt(List.of(userMessage),
// OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_VISION_PREVIEW.getValue()).build()));
String content = response.collectList().block().stream().collect(Collectors.joining());
logger.info("Response: {}", content);

View File

@@ -236,8 +236,8 @@ public interface ChatClient {
return this;
}
public ChatClientRequest functions(String... functions) {
this.functionNames.addAll(List.of(functions));
public ChatClientRequest functions(String... functionBeanNames) {
this.functionNames.addAll(List.of(functionBeanNames));
return this;
}
@@ -285,8 +285,7 @@ public interface ChatClient {
}
public <T> T single(ParameterizedTypeReference<T> type) {
return doSingleWithBeanOutputConverter(new BeanOutputConverter<T>(new ParameterizedTypeReference<>() {
}));
return doSingleWithBeanOutputConverter(new BeanOutputConverter<T>(type));
}
private <T> T doSingleWithBeanOutputConverter(BeanOutputConverter<T> boc) {
@@ -391,33 +390,6 @@ public interface ChatClient {
this.request = request;
}
// public <T> Flux<T> single(ParameterizedTypeReference<T> t) {
// return doSingleWithBeanOutputConverter(new BeanOutputConverter<T>(new
// ParameterizedTypeReference<>() {
// }));
// }
// private <T> Flux<T> doSingleWithBeanOutputConverter(BeanOutputConverter<T>
// boc) {
// var processedUserText = this.request.userText + System.lineSeparator() +
// System.lineSeparator()
// + "{format}";
// var chatResponse = doGetChatResponse(processedUserText, boc.getFormat());
// var stringResponse = chatResponse.getResult().getOutput().getContent();
// return boc.convert(stringResponse);
// }
// public <T> Flux<T> single(Class<T> clzz) {
// Assert.notNull(clzz, "the class must be non-null");
// var boc = new BeanOutputConverter<T>(clzz);
// return doSingleWithBeanOutputConverter(boc);
// }
// private Flux<ChatResponse> doGetFluxChatResponse(String processedUserText)
// {
// return this.doGetFluxChatResponse(processedUserText, "");
// }
private Flux<ChatResponse> doGetFluxChatResponse(String processedUserText) {
Map<String, Object> userParams = new HashMap<>(this.request.userParams);
@@ -475,16 +447,6 @@ public interface ChatClient {
}).filter(v -> StringUtils.hasText(v));
}
// @SuppressWarnings("unused")
// public <T> Collection<T> list(Class<T> clzz) {
// return single(new ParameterizedTypeReference<List<T>>() {
// });
// }
// public <T> Collection<T> list(ParameterizedTypeReference<List<T>> ptr) {
// return single(ptr);
// }
}
public CallResponseSpec call() {

View File

@@ -20,6 +20,7 @@ import java.net.MalformedURLException;
import java.net.URL;
import java.util.List;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.ArgumentCaptor;
@@ -48,13 +49,45 @@ public class ChatClientTest {
@Captor
ArgumentCaptor<Prompt> promptCaptor;
@BeforeEach
public void beforeAll() {
when(modelCaller.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
}
@Test
public void call() throws MalformedURLException {
public void simpleUserPrompt() throws MalformedURLException {
assertThat(ChatClient.builder(modelCaller).build().prompt().user("User prompt").call().content())
.isEqualTo("response");
Message userMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(userMessage.getContent()).isEqualTo("User prompt");
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
}
@Test
public void simpleSystemPrompt() throws MalformedURLException {
String response = ChatClient.builder(modelCaller).build().prompt().system("System prompt").call().content();
assertThat(response).isEqualTo("response");
assertThat(promptCaptor.getValue().getInstructions()).hasSize(2);
Message systemMessage = promptCaptor.getValue().getInstructions().get(0);
assertThat(systemMessage.getContent()).isEqualTo("System prompt");
assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM);
// Is this expected?
Message userMessage = promptCaptor.getValue().getInstructions().get(1);
assertThat(userMessage.getContent()).isEqualTo("");
assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER);
}
@Test
public void complexCall() throws MalformedURLException {
var options = FunctionCallingOptions.builder().build();
when(modelCaller.getDefaultOptions()).thenReturn(options);
when(modelCaller.call(promptCaptor.capture()))
.thenReturn(new ChatResponse(List.of(new Generation("response"))));
var url = new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png");

View File

@@ -28,6 +28,7 @@ import reactor.core.publisher.Flux;
import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.chat.messages.AssistantMessage;
@@ -87,18 +88,17 @@ class FunctionCallbackWithPlainFunctionBeanIT {
void functionCallWithPortableFunctionCallingOptions() {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> {
OpenAiModelCaller chatClient = context.getBean(OpenAiModelCaller.class);
OpenAiModelCaller caller = context.getBean(OpenAiModelCaller.class);
// Test weatherFunction
UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?");
// @formatter:off
String content = ChatClient.builder(caller).build().prompt()
.functions("weatherFunction")
.user("What's the weather like in San Francisco, Tokyo, and Paris?")
.stream().content()
.collectList().block().stream().collect(Collectors.joining());
// @formatter:on
PortableFunctionCallingOptions functionOptions = FunctionCallingOptions.builder()
.withFunction("weatherFunction")
.build();
ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), functionOptions));
logger.info("Response: {}", response);
logger.info("Response: {}", content);
});
}

View File

@@ -81,7 +81,7 @@ public class FunctionCallbackWrapper2IT {
// @formatter:off
String content = ChatClient.builder(caller).build().prompt()
.functions("WeatherInfo")
.user(u -> u.text("What's the weather like in San Francisco, Tokyo, and Paris?"))
.user("What's the weather like in San Francisco, Tokyo, and Paris?")
.stream().content()
.collectList().block().stream().collect(Collectors.joining());
// @formatter:on