From 545bc892b88ed90abe5c1443f72387749c4986d5 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Tue, 21 May 2024 23:13:33 +0200 Subject: [PATCH] Fix ChatClient list() convertions --- .../ai/openai/chat/OpenAiChatClientIT.java | 204 +++--------------- .../springframework/ai/chat/ChatClient.java | 44 +--- .../ai/chat/ChatClientTest.java | 39 +++- ...nctionCallbackWithPlainFunctionBeanIT.java | 20 +- .../tool/FunctionCallbackWrapper2IT.java | 2 +- 5 files changed, 85 insertions(+), 224 deletions(-) diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java index a0d8b8a70..100bfb566 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java @@ -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 list = ChatClient.builder(modelCaller).build().prompt() + Collection 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 list = - // outputConverter.convert(generation.getOutput().getContent()); - assertThat(list).hasSize(5); + // @formatter:off + List actorsFilms = ChatClient.builder(modelCaller).build().prompt() + .user("Generate the filmography of 5 movies for Tom Hanks and Bill Murray.") + .call() + .single(new ParameterizedTypeReference>() { + }); + // @formatter:on + + logger.info("" + actorsFilms); + assertThat(actorsFilms).hasSize(2); + + } + + @Test + void listOutputConverter3() { + + // @formatter:off + Collection 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 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>() { + .call() + .single(new ParameterizedTypeReference>() { }); // @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 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 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 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 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 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 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 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 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); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java index 655ff2336..2c304897f 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java @@ -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 single(ParameterizedTypeReference type) { - return doSingleWithBeanOutputConverter(new BeanOutputConverter(new ParameterizedTypeReference<>() { - })); + return doSingleWithBeanOutputConverter(new BeanOutputConverter(type)); } private T doSingleWithBeanOutputConverter(BeanOutputConverter boc) { @@ -391,33 +390,6 @@ public interface ChatClient { this.request = request; } - // public Flux single(ParameterizedTypeReference t) { - // return doSingleWithBeanOutputConverter(new BeanOutputConverter(new - // ParameterizedTypeReference<>() { - // })); - // } - - // private Flux doSingleWithBeanOutputConverter(BeanOutputConverter - // 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 Flux single(Class clzz) { - // Assert.notNull(clzz, "the class must be non-null"); - // var boc = new BeanOutputConverter(clzz); - // return doSingleWithBeanOutputConverter(boc); - // } - - // private Flux doGetFluxChatResponse(String processedUserText) - // { - // return this.doGetFluxChatResponse(processedUserText, ""); - // } - private Flux doGetFluxChatResponse(String processedUserText) { Map userParams = new HashMap<>(this.request.userParams); @@ -475,16 +447,6 @@ public interface ChatClient { }).filter(v -> StringUtils.hasText(v)); } - // @SuppressWarnings("unused") - // public Collection list(Class clzz) { - // return single(new ParameterizedTypeReference>() { - // }); - // } - - // public Collection list(ParameterizedTypeReference> ptr) { - // return single(ptr); - // } - } public CallResponseSpec call() { diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java index 2f2b2833a..ec9f17211 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java @@ -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 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"); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java index 2af30a86c..6f1ceb130 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java @@ -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); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java index 1608ca599..56b3df102 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java @@ -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