From 697a04b6d0263879aa390c5bda372a30f987155f Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Tue, 16 Jul 2024 17:55:24 +0200 Subject: [PATCH] Moonshot: No Function Calling support verificaiton Update the docs and tests to verify that currently Moonshot ChatModel does not support Function Calling. Related to #1058 --- .../ai/moonshot/MoonshotChatModel.java | 14 +-- .../MoonshotChatModelFunctionCallingIT.java | 111 ++++++++++++++++++ .../ai/moonshot/chat/MoonshotChatModelIT.java | 11 +- .../src/main/antora/modules/ROOT/nav.adoc | 2 +- .../ROOT/pages/api/chat/moonshot-chat.adoc | 2 +- 5 files changed, 125 insertions(+), 15 deletions(-) create mode 100644 models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/chat/MoonshotChatModelFunctionCallingIT.java diff --git a/models/spring-ai-moonshot/src/main/java/org/springframework/ai/moonshot/MoonshotChatModel.java b/models/spring-ai-moonshot/src/main/java/org/springframework/ai/moonshot/MoonshotChatModel.java index cd0a06804..373914494 100644 --- a/models/spring-ai-moonshot/src/main/java/org/springframework/ai/moonshot/MoonshotChatModel.java +++ b/models/spring-ai-moonshot/src/main/java/org/springframework/ai/moonshot/MoonshotChatModel.java @@ -15,13 +15,17 @@ */ package org.springframework.ai.moonshot; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; + import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.Generation; -import org.springframework.ai.chat.model.StreamingChatModel; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; @@ -34,17 +38,13 @@ import org.springframework.ai.retry.RetryUtils; import org.springframework.http.ResponseEntity; import org.springframework.retry.support.RetryTemplate; import org.springframework.util.Assert; -import reactor.core.publisher.Flux; -import java.util.HashMap; -import java.util.List; -import java.util.Map; -import java.util.concurrent.ConcurrentHashMap; +import reactor.core.publisher.Flux; /** * @author Geng Rong */ -public class MoonshotChatModel implements ChatModel, StreamingChatModel { +public class MoonshotChatModel implements ChatModel { private static final Logger logger = LoggerFactory.getLogger(MoonshotChatModel.class); diff --git a/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/chat/MoonshotChatModelFunctionCallingIT.java b/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/chat/MoonshotChatModelFunctionCallingIT.java new file mode 100644 index 000000000..f1a73b926 --- /dev/null +++ b/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/chat/MoonshotChatModelFunctionCallingIT.java @@ -0,0 +1,111 @@ +/* + * 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.moonshot.chat; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.ArrayList; +import java.util.List; +import java.util.stream.Collectors; + +import org.junit.jupiter.api.Disabled; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.model.Generation; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.moonshot.MoonshotChatOptions; +import org.springframework.ai.moonshot.MoonshotTestConfiguration; +import org.springframework.ai.moonshot.api.MockWeatherService; +import org.springframework.ai.moonshot.api.MoonshotApi; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; + +import reactor.core.publisher.Flux; + +@Disabled("Currently, the Moonshot Chat Model doesn't support function calling.") +@SpringBootTest(classes = MoonshotTestConfiguration.class) +@EnabledIfEnvironmentVariable(named = "MOONSHOT_API_KEY", matches = ".+") +class MoonshotChatModelFunctionCallingIT { + + private static final Logger logger = LoggerFactory.getLogger(MoonshotChatModelFunctionCallingIT.class); + + @Autowired + ChatModel chatModel; + + @Test + void functionCallTest() { + + UserMessage userMessage = new UserMessage( + "What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius."); + + List messages = new ArrayList<>(List.of(userMessage)); + + var promptOptions = MoonshotChatOptions.builder() + .withModel(MoonshotApi.ChatModel.MOONSHOT_V1_8K.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 = chatModel.call(new Prompt(messages, promptOptions)); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); + } + + @Test + void streamFunctionCallTest() { + + UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + + List messages = new ArrayList<>(List.of(userMessage)); + + var promptOptions = MoonshotChatOptions.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 = chatModel.stream(new Prompt(messages, promptOptions)); + + String content = response.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + logger.info("Response: {}", content); + + assertThat(content).contains("30", "10", "15"); + } + +} \ No newline at end of file diff --git a/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/chat/MoonshotChatModelIT.java b/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/chat/MoonshotChatModelIT.java index 9b7d10369..cbda0059d 100644 --- a/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/chat/MoonshotChatModelIT.java +++ b/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/chat/MoonshotChatModelIT.java @@ -75,7 +75,6 @@ public class MoonshotChatModelIT { ChatResponse response = chatModel.call(prompt); assertThat(response.getResults()).hasSize(1); assertThat(response.getResults().get(0).getOutput().getContent()).contains("Blackbeard"); - // needs fine tuning... evaluateQuestionAndAnswer(request, response, false); } @Test @@ -93,7 +92,7 @@ public class MoonshotChatModelIT { Prompt prompt = new Prompt(promptTemplate.createMessage()); Generation generation = this.chatModel.call(prompt).getResult(); - List list = outputConverter.parse(generation.getOutput().getContent()); + List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); } @@ -118,7 +117,7 @@ public class MoonshotChatModelIT { Prompt prompt = new Prompt(promptTemplate.createMessage()); Generation generation = chatModel.call(prompt).getResult(); - Map result = outputConverter.parse(generation.getOutput().getContent()); + Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); } @@ -137,7 +136,7 @@ public class MoonshotChatModelIT { Prompt prompt = new Prompt(promptTemplate.createMessage()); Generation generation = chatModel.call(prompt).getResult(); - ActorsFilms actorsFilms = outputConverter.parse(generation.getOutput().getContent()); + ActorsFilms actorsFilms = outputConverter.convert(generation.getOutput().getContent()); } @@ -165,7 +164,7 @@ public class MoonshotChatModelIT { Prompt prompt = new Prompt(promptTemplate.createMessage()); Generation generation = chatModel.call(prompt).getResult(); - ActorsFilmsRecord actorsFilms = outputConverter.parse(generation.getOutput().getContent()); + ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); logger.info("" + actorsFilms); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); assertThat(actorsFilms.movies()).hasSize(5); @@ -196,7 +195,7 @@ public class MoonshotChatModelIT { .map(AssistantMessage::getContent) .collect(Collectors.joining()); - ActorsFilmsRecord actorsFilms = outputConverter.parse(generationTextFromStream); + ActorsFilmsRecord actorsFilms = outputConverter.convert(generationTextFromStream); logger.info("" + actorsFilms); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); assertThat(actorsFilms.movies()).hasSize(5); diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/nav.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/nav.adoc index f364fa602..3901d0617 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/nav.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/nav.adoc @@ -26,7 +26,7 @@ *** xref:api/chat/minimax-chat.adoc[MiniMax] **** xref:api/chat/functions/minimax-chat-functions.adoc[Function Calling] *** xref:api/chat/moonshot-chat.adoc[Moonshot AI] -**** xref:api/chat/functions/moonshot-chat-functions.adoc[Function Calling] +//// **** xref:api/chat/functions/moonshot-chat-functions.adoc[Function Calling] *** xref:api/chat/ollama-chat.adoc[Ollama] *** xref:api/chat/openai-chat.adoc[OpenAI] **** xref:api/chat/functions/openai-chat-functions.adoc[Function Calling] diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/moonshot-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/moonshot-chat.adoc index 284b6a252..3b9383f70 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/moonshot-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/moonshot-chat.adoc @@ -247,4 +247,4 @@ Follow the https://github.com/spring-projects/spring-ai/blob/main/models/spring- ==== MoonshotApi Samples * The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/api/MoonshotApiIT.java[MoonshotApiIT.java] test provides some general examples how to use the lightweight library. -* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/api/MoonshotApiToolFunctionCallIT.java.java[MoonshotApiToolFunctionCallIT.java] test shows how to use the low-level API to call tool functions. \ No newline at end of file +* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-moonshot/src/test/java/org/springframework/ai/moonshot/api/MoonshotApiToolFunctionCallIT.java[MoonshotApiToolFunctionCallIT.java] test shows how to use the low-level API to call tool functions. \ No newline at end of file