From 4fac212b0db8af11cf8103aae5f27d78382ba5c6 Mon Sep 17 00:00:00 2001 From: GR Date: Fri, 23 Aug 2024 22:30:48 +0800 Subject: [PATCH] Add web search capability to MiniMax model Implement web search functionality for the MiniMax model. Includes unit tests This enhancement expands the model's ability to access and utilize current information from the internet. Resolves #1245 --- .../ai/minimax/MiniMaxChatModel.java | 3 ++- .../ai/minimax/api/MiniMaxApi.java | 8 +++++- .../MiniMaxStreamFunctionCallingHelper.java | 18 ++++++++++++- .../api/MiniMaxApiToolFunctionCallIT.java | 27 +++++++++++++++++++ .../ai/minimax/api/MiniMaxRetryTests.java | 2 +- 5 files changed, 54 insertions(+), 4 deletions(-) diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java index 257ea0264..288e5d7c1 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java @@ -17,6 +17,7 @@ package org.springframework.ai.minimax; import org.slf4j.Logger; import org.slf4j.LoggerFactory; + import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.messages.ToolResponseMessage; @@ -302,7 +303,7 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod if (delta == null) { delta = new ChatCompletionMessage("", Role.ASSISTANT); } - return new ChatCompletion.Choice(cc.finishReason(), cc.index(), delta, cc.logprobs()); + return new ChatCompletion.Choice(cc.finishReason(), cc.index(), delta, null, cc.logprobs()); }).toList(); return new ChatCompletion(chunk.id(), choices, chunk.created(), chunk.model(), chunk.systemFingerprint(), diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxApi.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxApi.java index 981e3335b..af040f853 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxApi.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxApi.java @@ -160,6 +160,10 @@ public class MiniMaxApi { this(Type.FUNCTION, function); } + public static FunctionTool webSearchFunctionTool() { + return new FunctionTool(Type.WEB_SEARCH, null); + } + /** * Create a tool of type 'function' and the given function definition. */ @@ -167,7 +171,8 @@ public class MiniMaxApi { /** * Function tool type. */ - @JsonProperty("function") FUNCTION + @JsonProperty("function") FUNCTION, + @JsonProperty("web_search") WEB_SEARCH } /** @@ -561,6 +566,7 @@ public class MiniMaxApi { @JsonProperty("finish_reason") ChatCompletionFinishReason finishReason, @JsonProperty("index") Integer index, @JsonProperty("message") ChatCompletionMessage message, + @JsonProperty("messages") List messages, @JsonProperty("logprobs") LogProbs logprobs) { } diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxStreamFunctionCallingHelper.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxStreamFunctionCallingHelper.java index 06454a348..82b2eca12 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxStreamFunctionCallingHelper.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxStreamFunctionCallingHelper.java @@ -170,7 +170,23 @@ public class MiniMaxStreamFunctionCallingHelper { if (choice == null || choice.delta() == null) { return false; } - return choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; + return choice.finishReason() == MiniMaxApi.ChatCompletionFinishReason.TOOL_CALLS; + } + + /** + * Convert the ChatCompletionChunk into a ChatCompletion. The Usage is set to null. + * @param chunk the ChatCompletionChunk to convert + * @return the ChatCompletion + */ + public MiniMaxApi.ChatCompletion chunkToChatCompletion(MiniMaxApi.ChatCompletionChunk chunk) { + List choices = chunk.choices() + .stream() + .map(chunkChoice -> new MiniMaxApi.ChatCompletion.Choice(chunkChoice.finishReason(), chunkChoice.index(), + chunkChoice.delta(), null, chunkChoice.logprobs())) + .toList(); + + return new MiniMaxApi.ChatCompletion(chunk.id(), choices, chunk.created(), chunk.model(), + chunk.systemFingerprint(), "chat.completion", null, null); } } diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxApiToolFunctionCallIT.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxApiToolFunctionCallIT.java index 77abfd0ce..7f599b8e0 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxApiToolFunctionCallIT.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxApiToolFunctionCallIT.java @@ -131,6 +131,33 @@ public class MiniMaxApiToolFunctionCallIT { .containsAnyOf("°C", "Celsius"); } + @SuppressWarnings("null") + @Test + public void webSearchToolFunctionCall() { + + var message = new ChatCompletionMessage( + "How many gold medals has the United States won in total at the 2024 Olympics?", Role.USER); + + var functionTool = MiniMaxApi.FunctionTool.webSearchFunctionTool(); + + List messages = new ArrayList<>(List.of(message)); + + ChatCompletionRequest chatCompletionRequest = new ChatCompletionRequest(messages, + org.springframework.ai.minimax.api.MiniMaxApi.ChatModel.ABAB_6_5_S_Chat.getValue(), + List.of(functionTool), ToolChoiceBuilder.AUTO); + + ResponseEntity chatCompletion = miniMaxApi.chatCompletionEntity(chatCompletionRequest); + + assertThat(chatCompletion.getBody()).isNotNull(); + assertThat(chatCompletion.getBody().choices()).isNotEmpty(); + + List responseMessages = chatCompletion.getBody().choices().get(0).messages(); + ChatCompletionMessage assistantMessage = responseMessages.get(responseMessages.size() - 1); + + assertThat(assistantMessage.role()).isEqualTo(Role.ASSISTANT); + assertThat(assistantMessage.content()).contains("40"); + } + private static T fromJson(String json, Class targetClass) { try { return new ObjectMapper().readValue(json, targetClass); diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java index 493874de5..0978b4a8b 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java @@ -102,7 +102,7 @@ public class MiniMaxRetryTests { public void miniMaxChatTransientError() { var choice = new ChatCompletion.Choice(ChatCompletionFinishReason.STOP, 0, - new ChatCompletionMessage("Response", Role.ASSISTANT), null); + new ChatCompletionMessage("Response", Role.ASSISTANT), null, null); ChatCompletion expectedChatCompletion = new ChatCompletion("id", List.of(choice), 666l, "model", null, null, null, new MiniMaxApi.Usage(10, 10, 10));