Refactor toolcalling support for Minimax

- Update Minimax Chat Model to use ToolCalling Manager and ToolExecutionEligibilityPredicate
 - Update Minimax ChatOptions to implement ToolCallingChatOptions
 - Update Autoconfiguration for MinimaxChat model to use ToolCallingAutoconfiguration
 - Update tests

Signed-off-by: Ilayaperumal Gopinathan <ilayaperumal.gopinathan@broadcom.com>
This commit is contained in:
Ilayaperumal Gopinathan
2025-04-25 11:38:47 +01:00
parent 98052624b1
commit 49e1dd10a9
10 changed files with 322 additions and 288 deletions

View File

@@ -16,8 +16,6 @@
package org.springframework.ai.model.minimax.autoconfigure;
import java.util.List;
import io.micrometer.observation.ObservationRegistry;
import org.springframework.ai.chat.observation.ChatModelObservationConvention;
@@ -25,18 +23,19 @@ import org.springframework.ai.minimax.MiniMaxChatModel;
import org.springframework.ai.minimax.api.MiniMaxApi;
import org.springframework.ai.model.SpringAIModelProperties;
import org.springframework.ai.model.SpringAIModels;
import org.springframework.ai.model.function.DefaultFunctionCallbackResolver;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackResolver;
import org.springframework.ai.model.tool.DefaultToolExecutionEligibilityPredicate;
import org.springframework.ai.model.tool.ToolCallingManager;
import org.springframework.ai.model.tool.ToolExecutionEligibilityPredicate;
import org.springframework.ai.model.tool.autoconfigure.ToolCallingAutoConfiguration;
import org.springframework.ai.retry.autoconfigure.SpringAiRetryAutoConfiguration;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.ImportAutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.ApplicationContext;
import org.springframework.context.annotation.Bean;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
@@ -50,28 +49,32 @@ import org.springframework.web.client.RestClient;
* @author Geng Rong
* @author Ilayaperumal Gopinathan
*/
@AutoConfiguration(after = { RestClientAutoConfiguration.class, SpringAiRetryAutoConfiguration.class })
@AutoConfiguration(after = { RestClientAutoConfiguration.class, SpringAiRetryAutoConfiguration.class,
ToolCallingAutoConfiguration.class })
@ConditionalOnClass(MiniMaxApi.class)
@EnableConfigurationProperties({ MiniMaxConnectionProperties.class, MiniMaxChatProperties.class })
@ConditionalOnProperty(name = SpringAIModelProperties.CHAT_MODEL, havingValue = SpringAIModels.MINIMAX,
matchIfMissing = true)
@ImportAutoConfiguration(classes = { SpringAiRetryAutoConfiguration.class, RestClientAutoConfiguration.class,
ToolCallingAutoConfiguration.class })
public class MiniMaxChatAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public MiniMaxChatModel miniMaxChatModel(MiniMaxConnectionProperties commonProperties,
MiniMaxChatProperties chatProperties, ObjectProvider<RestClient.Builder> restClientBuilderProvider,
List<FunctionCallback> toolFunctionCallbacks, FunctionCallbackResolver functionCallbackResolver,
RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler,
ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<ChatModelObservationConvention> observationConvention) {
ToolCallingManager toolCallingManager, RetryTemplate retryTemplate,
ResponseErrorHandler responseErrorHandler, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<ChatModelObservationConvention> observationConvention,
ObjectProvider<ToolExecutionEligibilityPredicate> openAiToolExecutionEligibilityPredicate) {
var miniMaxApi = miniMaxApi(chatProperties.getBaseUrl(), commonProperties.getBaseUrl(),
chatProperties.getApiKey(), commonProperties.getApiKey(),
restClientBuilderProvider.getIfAvailable(RestClient::builder), responseErrorHandler);
var chatModel = new MiniMaxChatModel(miniMaxApi, chatProperties.getOptions(), functionCallbackResolver,
toolFunctionCallbacks, retryTemplate, observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP));
var chatModel = new MiniMaxChatModel(miniMaxApi, chatProperties.getOptions(), toolCallingManager, retryTemplate,
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
openAiToolExecutionEligibilityPredicate.getIfUnique(DefaultToolExecutionEligibilityPredicate::new));
observationConvention.ifAvailable(chatModel::setObservationConvention);
return chatModel;
@@ -89,12 +92,4 @@ public class MiniMaxChatAutoConfiguration {
return new MiniMaxApi(resolvedBaseUrl, resolvedApiKey, restClientBuilder, responseErrorHandler);
}
@Bean
@ConditionalOnMissingBean
public FunctionCallbackResolver springAiFunctionManager(ApplicationContext context) {
DefaultFunctionCallbackResolver manager = new DefaultFunctionCallbackResolver();
manager.setApplicationContext(context);
return manager;
}
}

View File

@@ -34,6 +34,7 @@ import org.springframework.ai.minimax.MiniMaxChatModel;
import org.springframework.ai.minimax.MiniMaxChatOptions;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.retry.autoconfigure.SpringAiRetryAutoConfiguration;
import org.springframework.ai.tool.function.FunctionToolCallback;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
@@ -63,8 +64,7 @@ public class FunctionCallbackInPromptIT {
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
var promptOptions = MiniMaxChatOptions.builder()
.functionCallbacks(List.of(FunctionCallback.builder()
.function("CurrentWeatherService", new MockWeatherService())
.toolCallbacks(List.of(FunctionToolCallback.builder("CurrentWeatherService", new MockWeatherService())
.description("Get the weather in location")
.inputType(MockWeatherService.Request.class)
.build()))
@@ -89,8 +89,7 @@ public class FunctionCallbackInPromptIT {
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
var promptOptions = MiniMaxChatOptions.builder()
.functionCallbacks(List.of(FunctionCallback.builder()
.function("CurrentWeatherService", new MockWeatherService())
.toolCallbacks(List.of(FunctionToolCallback.builder("CurrentWeatherService", new MockWeatherService())
.description("Get the weather in location")
.inputType(MockWeatherService.Request.class)
.build()))

View File

@@ -34,6 +34,7 @@ import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.minimax.MiniMaxChatModel;
import org.springframework.ai.minimax.MiniMaxChatOptions;
import org.springframework.ai.model.function.FunctionCallingOptions;
import org.springframework.ai.model.tool.ToolCallingChatOptions;
import org.springframework.ai.retry.autoconfigure.SpringAiRetryAutoConfiguration;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
@@ -69,8 +70,8 @@ class FunctionCallbackWithPlainFunctionBeanIT {
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
ChatResponse response = chatModel.call(
new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().function("weatherFunction").build()));
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
MiniMaxChatOptions.builder().toolNames("weatherFunction").build()));
logger.info("Response: {}", response);
@@ -78,7 +79,7 @@ class FunctionCallbackWithPlainFunctionBeanIT {
// Test weatherFunctionTwo
response = chatModel.call(new Prompt(List.of(userMessage),
MiniMaxChatOptions.builder().function("weatherFunctionTwo").build()));
MiniMaxChatOptions.builder().toolNames("weatherFunctionTwo").build()));
logger.info("Response: {}", response);
@@ -97,8 +98,8 @@ class FunctionCallbackWithPlainFunctionBeanIT {
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
FunctionCallingOptions functionOptions = FunctionCallingOptions.builder()
.function("weatherFunction")
ToolCallingChatOptions functionOptions = ToolCallingChatOptions.builder()
.toolNames("weatherFunction")
.build();
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions));
@@ -118,8 +119,8 @@ class FunctionCallbackWithPlainFunctionBeanIT {
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
Flux<ChatResponse> response = chatModel.stream(
new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().function("weatherFunction").build()));
Flux<ChatResponse> response = chatModel.stream(new Prompt(List.of(userMessage),
MiniMaxChatOptions.builder().toolNames("weatherFunction").build()));
String content = response.collectList()
.block()
@@ -137,7 +138,7 @@ class FunctionCallbackWithPlainFunctionBeanIT {
// Test weatherFunctionTwo
response = chatModel.stream(new Prompt(List.of(userMessage),
MiniMaxChatOptions.builder().function("weatherFunctionTwo").build()));
MiniMaxChatOptions.builder().toolNames("weatherFunctionTwo").build()));
content = response.collectList()
.block()

View File

@@ -33,7 +33,9 @@ import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.minimax.MiniMaxChatModel;
import org.springframework.ai.minimax.MiniMaxChatOptions;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.tool.autoconfigure.ToolCallingAutoConfiguration;
import org.springframework.ai.retry.autoconfigure.SpringAiRetryAutoConfiguration;
import org.springframework.ai.tool.function.FunctionToolCallback;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
@@ -52,8 +54,9 @@ public class MiniMaxFunctionCallbackIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.minimax.apiKey=" + System.getenv("MINIMAX_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, MiniMaxChatAutoConfiguration.class))
.withConfiguration(
AutoConfigurations.of(SpringAiRetryAutoConfiguration.class, RestClientAutoConfiguration.class,
MiniMaxChatAutoConfiguration.class, ToolCallingAutoConfiguration.class))
.withUserConfiguration(Config.class);
@Test
@@ -66,7 +69,7 @@ public class MiniMaxFunctionCallbackIT {
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
ChatResponse response = chatModel
.call(new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().function("WeatherInfo").build()));
.call(new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().toolNames("WeatherInfo").build()));
logger.info("Response: {}", response);
@@ -84,8 +87,8 @@ public class MiniMaxFunctionCallbackIT {
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? Return the temperature in Celsius.");
Flux<ChatResponse> response = chatModel
.stream(new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().function("WeatherInfo").build()));
Flux<ChatResponse> response = chatModel.stream(
new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().toolNames("WeatherInfo").build()));
String content = response.collectList()
.block()
@@ -108,10 +111,9 @@ public class MiniMaxFunctionCallbackIT {
static class Config {
@Bean
public FunctionCallback weatherFunctionInfo() {
public FunctionToolCallback<MockWeatherService.Request, MockWeatherService.Response> weatherFunctionInfo() {
return FunctionCallback.builder()
.function("WeatherInfo", new MockWeatherService())
return FunctionToolCallback.builder("WeatherInfo", new MockWeatherService())
.description("Get the weather in location")
.inputType(MockWeatherService.Request.class)
.build();

View File

@@ -17,10 +17,8 @@
package org.springframework.ai.minimax;
import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
import io.micrometer.observation.Observation;
@@ -39,7 +37,6 @@ import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.DefaultUsage;
import org.springframework.ai.chat.metadata.EmptyUsage;
import org.springframework.ai.chat.model.AbstractToolCallSupport;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
@@ -63,10 +60,13 @@ import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletionMessage.ToolC
import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletionRequest;
import org.springframework.ai.minimax.api.MiniMaxApiConstants;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackResolver;
import org.springframework.ai.model.function.FunctionCallingOptions;
import org.springframework.ai.model.tool.DefaultToolExecutionEligibilityPredicate;
import org.springframework.ai.model.tool.ToolCallingChatOptions;
import org.springframework.ai.model.tool.ToolCallingManager;
import org.springframework.ai.model.tool.ToolExecutionEligibilityPredicate;
import org.springframework.ai.model.tool.ToolExecutionResult;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.ai.tool.definition.ToolDefinition;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
@@ -78,12 +78,13 @@ import org.springframework.util.CollectionUtils;
*
* @author Geng Rong
* @author Alexandros Pappas
* @author Ilayaperumal Gopinathan
* @see ChatModel
* @see StreamingChatModel
* @see MiniMaxApi
* @since 1.0.0 M1
*/
public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatModel, StreamingChatModel {
public class MiniMaxChatModel implements ChatModel {
private static final Logger logger = LoggerFactory.getLogger(MiniMaxChatModel.class);
@@ -109,6 +110,17 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
*/
private final ObservationRegistry observationRegistry;
/**
* The tool calling manager.
*/
private final ToolCallingManager toolCallingManager;
/**
* The tool execution eligibility predicate used to determine if a tool can be
* executed.
*/
private final ToolExecutionEligibilityPredicate toolExecutionEligibilityPredicate;
/**
* Conventions to use for generating observations.
*/
@@ -131,7 +143,7 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
* @param options The MiniMaxChatOptions to configure the chat model.
*/
public MiniMaxChatModel(MiniMaxApi miniMaxApi, MiniMaxChatOptions options) {
this(miniMaxApi, options, null, RetryUtils.DEFAULT_RETRY_TEMPLATE);
this(miniMaxApi, options, ToolCallingManager.builder().build(), RetryUtils.DEFAULT_RETRY_TEMPLATE);
}
/**
@@ -139,13 +151,24 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
* @param miniMaxApi The MiniMaxApi instance to be used for interacting with the
* MiniMax Chat API.
* @param options The MiniMaxChatOptions to configure the chat model.
* @param functionCallbackResolver The function callback resolver to resolve the
* function by its name.
* @param toolCallingManager The tool calling manager.
*/
public MiniMaxChatModel(MiniMaxApi miniMaxApi, MiniMaxChatOptions options, ToolCallingManager toolCallingManager) {
this(miniMaxApi, options, toolCallingManager, RetryUtils.DEFAULT_RETRY_TEMPLATE);
}
/**
* Initializes a new instance of the MiniMaxChatModel.
* @param miniMaxApi The MiniMaxApi instance to be used for interacting with the
* MiniMax Chat API.
* @param options The MiniMaxChatOptions to configure the chat model.
* @param toolCallingManager The tool calling manager.
* @param retryTemplate The retry template.
*/
public MiniMaxChatModel(MiniMaxApi miniMaxApi, MiniMaxChatOptions options,
FunctionCallbackResolver functionCallbackResolver, RetryTemplate retryTemplate) {
this(miniMaxApi, options, functionCallbackResolver, List.of(), retryTemplate, ObservationRegistry.NOOP);
public MiniMaxChatModel(MiniMaxApi miniMaxApi, MiniMaxChatOptions options, ToolCallingManager toolCallingManager,
RetryTemplate retryTemplate) {
this(miniMaxApi, options, toolCallingManager, retryTemplate, ObservationRegistry.NOOP,
new DefaultToolExecutionEligibilityPredicate());
}
/**
@@ -153,26 +176,25 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
* @param miniMaxApi The MiniMaxApi instance to be used for interacting with the
* MiniMax Chat API.
* @param options The MiniMaxChatOptions to configure the chat model.
* @param functionCallbackResolver The function callback resolver to resolve the
* function by its name.
* @param toolFunctionCallbacks The tool function callbacks.
* @param retryTemplate The retry template.
* @param observationRegistry The ObservationRegistry used for instrumentation.
* @param toolExecutionEligibilityPredicate The Tool
*/
public MiniMaxChatModel(MiniMaxApi miniMaxApi, MiniMaxChatOptions options,
FunctionCallbackResolver functionCallbackResolver, List<FunctionCallback> toolFunctionCallbacks,
RetryTemplate retryTemplate, ObservationRegistry observationRegistry) {
super(functionCallbackResolver, options, toolFunctionCallbacks);
public MiniMaxChatModel(MiniMaxApi miniMaxApi, MiniMaxChatOptions options, ToolCallingManager toolCallingManager,
RetryTemplate retryTemplate, ObservationRegistry observationRegistry,
ToolExecutionEligibilityPredicate toolExecutionEligibilityPredicate) {
Assert.notNull(miniMaxApi, "MiniMaxApi must not be null");
Assert.notNull(options, "Options must not be null");
Assert.notNull(toolCallingManager, "toolCallingManager cannot be null");
Assert.notNull(retryTemplate, "RetryTemplate must not be null");
Assert.isTrue(CollectionUtils.isEmpty(options.getFunctionCallbacks()),
"The default function callbacks must be set via the toolFunctionCallbacks constructor parameter");
Assert.notNull(observationRegistry, "ObservationRegistry must not be null");
Assert.notNull(toolExecutionEligibilityPredicate, "toolExecutionEligibilityPredicate cannot be null");
this.miniMaxApi = miniMaxApi;
this.defaultOptions = options;
this.toolCallingManager = toolCallingManager;
this.retryTemplate = retryTemplate;
this.observationRegistry = observationRegistry;
this.toolExecutionEligibilityPredicate = toolExecutionEligibilityPredicate;
}
private static Generation buildGeneration(Choice choice, Map<String, Object> metadata) {
@@ -211,12 +233,15 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
@Override
public ChatResponse call(Prompt prompt) {
ChatCompletionRequest request = createRequest(prompt, false);
// Before moving any further, build the final request Prompt,
// merging runtime and default options.
Prompt requestPrompt = buildRequestPrompt(prompt);
ChatCompletionRequest request = createRequest(requestPrompt, false);
ChatModelObservationContext observationContext = ChatModelObservationContext.builder()
.prompt(prompt)
.prompt(requestPrompt)
.provider(MiniMaxApiConstants.PROVIDER_NAME)
.requestOptions(buildRequestOptions(request))
.requestOptions(requestPrompt.getOptions())
.build();
ChatResponse response = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION
@@ -230,13 +255,13 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
var chatCompletion = completionEntity.getBody();
if (chatCompletion == null) {
logger.warn("No chat completion returned for prompt: {}", prompt);
logger.warn("No chat completion returned for prompt: {}", requestPrompt);
return new ChatResponse(List.of());
}
List<Choice> choices = chatCompletion.choices();
if (choices == null) {
logger.warn("No choices returned for prompt: {}, because: {}}", prompt,
logger.warn("No choices returned for prompt: {}, because: {}}", requestPrompt,
chatCompletion.baseResponse().message());
return new ChatResponse(List.of());
}
@@ -268,12 +293,19 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
return chatResponse;
});
if (!isProxyToolCalls(prompt, this.defaultOptions) && isToolCall(response,
Set.of(ChatCompletionFinishReason.TOOL_CALLS.name(), ChatCompletionFinishReason.STOP.name()))) {
var toolCallConversation = handleToolCalls(prompt, response);
// Recursively call the call method with the tool call message
// conversation that contains the call responses.
return this.call(new Prompt(toolCallConversation, prompt.getOptions()));
if (this.toolExecutionEligibilityPredicate.isToolExecutionRequired(requestPrompt.getOptions(), response)) {
var toolExecutionResult = this.toolCallingManager.executeToolCalls(requestPrompt, response);
if (toolExecutionResult.returnDirect()) {
// Return tool execution result directly to the client.
return ChatResponse.builder()
.from(response)
.generations(ToolExecutionResult.buildGenerations(toolExecutionResult))
.build();
}
else {
// Send the tool execution result back to the model.
return this.call(new Prompt(toolExecutionResult.conversationHistory(), requestPrompt.getOptions()));
}
}
return response;
@@ -286,8 +318,11 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
@Override
public Flux<ChatResponse> stream(Prompt prompt) {
// Before moving any further, build the final request Prompt,
// merging runtime and default options.
Prompt requestPrompt = buildRequestPrompt(prompt);
return Flux.deferContextual(contextView -> {
ChatCompletionRequest request = createRequest(prompt, true);
ChatCompletionRequest request = createRequest(requestPrompt, true);
Flux<ChatCompletionChunk> completionChunks = this.retryTemplate
.execute(ctx -> this.miniMaxApi.chatCompletionStream(request));
@@ -297,9 +332,9 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
ConcurrentHashMap<String, String> roleMap = new ConcurrentHashMap<>();
final ChatModelObservationContext observationContext = ChatModelObservationContext.builder()
.prompt(prompt)
.prompt(requestPrompt)
.provider(MiniMaxApiConstants.PROVIDER_NAME)
.requestOptions(buildRequestOptions(request))
.requestOptions(requestPrompt.getOptions())
.build();
Observation observation = ChatModelObservationDocumentation.CHAT_MODEL_OPERATION.observation(
@@ -336,15 +371,21 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
}));
Flux<ChatResponse> flux = chatResponse.flatMap(response -> {
if (!isProxyToolCalls(prompt, this.defaultOptions) && isToolCall(response,
Set.of(ChatCompletionFinishReason.TOOL_CALLS.name(), ChatCompletionFinishReason.STOP.name()))) {
// FIXME: bounded elastic needs to be used since tool calling
// is currently only synchronous
if (this.toolExecutionEligibilityPredicate.isToolExecutionRequired(requestPrompt.getOptions(), response)) {
return Flux.defer(() -> {
var toolCallConversation = handleToolCalls(prompt, response);
// Recursively call the stream method with the tool call message
// conversation that contains the call responses.
return this.stream(new Prompt(toolCallConversation, prompt.getOptions()));
// FIXME: bounded elastic needs to be used since tool calling
// is currently only synchronous
var toolExecutionResult = this.toolCallingManager.executeToolCalls(requestPrompt, response);
if (toolExecutionResult.returnDirect()) {
// Return tool execution result directly to the client.
return Flux.just(ChatResponse.builder().from(response)
.generations(ToolExecutionResult.buildGenerations(toolExecutionResult))
.build());
}
else {
// Send the tool execution result back to the model.
return this.stream(new Prompt(toolExecutionResult.conversationHistory(), requestPrompt.getOptions()));
}
}).subscribeOn(Schedulers.boundedElastic());
}
return Flux.just(response);
@@ -358,37 +399,6 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
});
}
/**
* The MimiMax web search function tool type is 'web_search', so we need to filter out
* the tool calls whose type is not 'function'
* @param generation the generation to check
* @param toolCallFinishReasons the tool call finish reasons
* @return true if the generation is a tool call
*/
@Override
protected boolean isToolCall(Generation generation, Set<String> toolCallFinishReasons) {
if (!super.isToolCall(generation, toolCallFinishReasons)) {
return false;
}
return generation.getOutput()
.getToolCalls()
.stream()
.anyMatch(toolCall -> org.springframework.ai.minimax.api.MiniMaxApiConstants.TOOL_CALL_FUNCTION_TYPE
.equals(toolCall.type()));
}
private ChatOptions buildRequestOptions(ChatCompletionRequest request) {
return ChatOptions.builder()
.model(request.model())
.frequencyPenalty(request.frequencyPenalty())
.maxTokens(request.maxTokens())
.presencePenalty(request.presencePenalty())
.stopSequences(request.stop())
.temperature(request.temperature())
.topP(request.topP())
.build();
}
private ChatResponseMetadata from(ChatCompletion result) {
Assert.notNull(result, "MiniMax ChatCompletionResult must not be null");
return ChatResponseMetadata.builder()
@@ -440,6 +450,49 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
"chat.completion", null, null);
}
Prompt buildRequestPrompt(Prompt prompt) {
// Process runtime options
MiniMaxChatOptions runtimeOptions = null;
if (prompt.getOptions() != null) {
if (prompt.getOptions() instanceof ToolCallingChatOptions toolCallingChatOptions) {
runtimeOptions = ModelOptionsUtils.copyToTarget(toolCallingChatOptions, ToolCallingChatOptions.class,
MiniMaxChatOptions.class);
}
else {
runtimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class,
MiniMaxChatOptions.class);
}
}
// Define request options by merging runtime options and default options
MiniMaxChatOptions requestOptions = ModelOptionsUtils.merge(runtimeOptions, this.defaultOptions,
MiniMaxChatOptions.class);
// Merge @JsonIgnore-annotated options explicitly since they are ignored by
// Jackson, used by ModelOptionsUtils.
if (runtimeOptions != null) {
requestOptions.setInternalToolExecutionEnabled(
ModelOptionsUtils.mergeOption(runtimeOptions.getInternalToolExecutionEnabled(),
this.defaultOptions.getInternalToolExecutionEnabled()));
requestOptions.setToolNames(ToolCallingChatOptions.mergeToolNames(runtimeOptions.getToolNames(),
this.defaultOptions.getToolNames()));
requestOptions.setToolCallbacks(ToolCallingChatOptions.mergeToolCallbacks(runtimeOptions.getToolCallbacks(),
this.defaultOptions.getToolCallbacks()));
requestOptions.setToolContext(ToolCallingChatOptions.mergeToolContext(runtimeOptions.getToolContext(),
this.defaultOptions.getToolContext()));
}
else {
requestOptions.setInternalToolExecutionEnabled(this.defaultOptions.getInternalToolExecutionEnabled());
requestOptions.setToolNames(this.defaultOptions.getToolNames());
requestOptions.setToolCallbacks(this.defaultOptions.getToolCallbacks());
requestOptions.setToolContext(this.defaultOptions.getToolContext());
}
ToolCallingChatOptions.validateToolCallbacks(requestOptions.getToolCallbacks());
return new Prompt(prompt.getInstructions(), requestOptions);
}
/**
* Accessible for testing.
*/
@@ -481,46 +534,41 @@ public class MiniMaxChatModel extends AbstractToolCallSupport implements ChatMod
}).flatMap(List::stream).toList();
ChatCompletionRequest request = new ChatCompletionRequest(chatCompletionMessages, stream);
MiniMaxChatOptions requestOptions = (MiniMaxChatOptions) prompt.getOptions();
request = ModelOptionsUtils.merge(requestOptions, request, ChatCompletionRequest.class);
Set<String> enabledToolsToUse = new HashSet<>();
// Add the tool definitions to the request's tools parameter.
List<ToolDefinition> toolDefinitions = this.toolCallingManager.resolveToolDefinitions(requestOptions);
if (!CollectionUtils.isEmpty(toolDefinitions)) {
request = ModelOptionsUtils.merge(
MiniMaxChatOptions.builder().tools(this.getFunctionTools(toolDefinitions)).build(), request,
ChatCompletionRequest.class);
}
if (prompt.getOptions() != null) {
MiniMaxChatOptions updatedRuntimeOptions;
if (prompt.getOptions() instanceof FunctionCallingOptions functionCallingOptions) {
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(functionCallingOptions,
FunctionCallingOptions.class, MiniMaxChatOptions.class);
if (prompt.getOptions() instanceof ToolCallingChatOptions toolCallingChatOptions) {
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(toolCallingChatOptions,
ToolCallingChatOptions.class, MiniMaxChatOptions.class);
}
else {
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class,
MiniMaxChatOptions.class);
}
enabledToolsToUse.addAll(this.runtimeFunctionCallbackConfigurations(updatedRuntimeOptions));
request = ModelOptionsUtils.merge(updatedRuntimeOptions, request, ChatCompletionRequest.class);
}
if (!CollectionUtils.isEmpty(this.defaultOptions.getFunctions())) {
enabledToolsToUse.addAll(this.defaultOptions.getFunctions());
}
request = ModelOptionsUtils.merge(request, this.defaultOptions, ChatCompletionRequest.class);
if (!CollectionUtils.isEmpty(enabledToolsToUse)) {
request = ModelOptionsUtils.merge(
MiniMaxChatOptions.builder().tools(this.getFunctionTools(enabledToolsToUse)).build(), request,
ChatCompletionRequest.class);
}
return request;
}
private List<MiniMaxApi.FunctionTool> getFunctionTools(Set<String> functionNames) {
return this.resolveFunctionCallbacks(functionNames).stream().map(functionCallback -> {
var function = new MiniMaxApi.FunctionTool.Function(functionCallback.getDescription(),
functionCallback.getName(), functionCallback.getInputTypeSchema());
private List<MiniMaxApi.FunctionTool> getFunctionTools(List<ToolDefinition> toolDefinitions) {
return toolDefinitions.stream().map(toolDefinition -> {
var function = new MiniMaxApi.FunctionTool.Function(toolDefinition.description(), toolDefinition.name(),
toolDefinition.inputSchema());
return new MiniMaxApi.FunctionTool(function);
}).toList();
}

View File

@@ -17,6 +17,8 @@
package org.springframework.ai.minimax;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
@@ -29,8 +31,9 @@ import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.minimax.api.MiniMaxApi;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallingOptions;
import org.springframework.ai.model.tool.ToolCallingChatOptions;
import org.springframework.ai.tool.ToolCallback;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
/**
@@ -38,7 +41,6 @@ import org.springframework.util.Assert;
* MiniMax API. It provides methods to set and retrieve various options like model,
* frequency penalty, max tokens, etc.
*
* @see FunctionCallingOptions
* @see ChatOptions
* @author Geng Rong
* @author Thomas Vitale
@@ -46,7 +48,7 @@ import org.springframework.util.Assert;
* @since 1.0.0 M1
*/
@JsonInclude(Include.NON_NULL)
public class MiniMaxChatOptions implements FunctionCallingOptions {
public class MiniMaxChatOptions implements ToolCallingChatOptions {
// @formatter:off
/**
@@ -128,25 +130,28 @@ public class MiniMaxChatOptions implements FunctionCallingOptions {
* from the registry to be used by the ChatModel chat completion requests.
*/
@JsonIgnore
private List<FunctionCallback> functionCallbacks = new ArrayList<>();
private List<ToolCallback> toolCallbacks = new ArrayList<>();
/**
* List of functions, identified by their names, to configure for function calling in
* the chat completion requests.
* Functions with those names must exist in the functionCallbacks registry.
* The {@link #functionCallbacks} from the PromptOptions are automatically enabled for the duration of the prompt execution.
* The {@link #toolCallbacks} from the PromptOptions are automatically enabled for the duration of the prompt execution.
*
* Note that function enabled with the default options are enabled for all chat completion requests. This could impact the token count and the billing.
* If the functions is set in a prompt options, then the enabled functions are only active for the duration of this prompt execution.
*/
@JsonIgnore
private Set<String> functions = new HashSet<>();
private Set<String> toolNames = new HashSet<>();
@JsonIgnore
private Boolean proxyToolCalls;
private Map<String, Object> toolContext = new HashMap<>();
/**
* Whether to enable the tool execution lifecycle internally in ChatModel.
*/
@JsonIgnore
private Map<String, Object> toolContext;
private Boolean internalToolExecutionEnabled;
// @formatter:on
@@ -168,9 +173,9 @@ public class MiniMaxChatOptions implements FunctionCallingOptions {
.maskSensitiveInfo(fromOptions.getMaskSensitiveInfo())
.tools(fromOptions.getTools())
.toolChoice(fromOptions.getToolChoice())
.functionCallbacks(fromOptions.getFunctionCallbacks())
.functions(fromOptions.getFunctions())
.proxyToolCalls(fromOptions.getProxyToolCalls())
.toolCallbacks(fromOptions.getToolCallbacks())
.toolNames(fromOptions.getToolNames())
.internalToolExecutionEnabled(fromOptions.getInternalToolExecutionEnabled())
.toolContext(fromOptions.getToolContext())
.build();
}
@@ -296,25 +301,6 @@ public class MiniMaxChatOptions implements FunctionCallingOptions {
this.toolChoice = toolChoice;
}
@Override
public List<FunctionCallback> getFunctionCallbacks() {
return this.functionCallbacks;
}
@Override
public void setFunctionCallbacks(List<FunctionCallback> functionCallbacks) {
this.functionCallbacks = functionCallbacks;
}
@Override
public Set<String> getFunctions() {
return this.functions;
}
public void setFunctions(Set<String> functionNames) {
this.functions = functionNames;
}
@Override
@JsonIgnore
public Integer getTopK() {
@@ -322,12 +308,45 @@ public class MiniMaxChatOptions implements FunctionCallingOptions {
}
@Override
public Boolean getProxyToolCalls() {
return this.proxyToolCalls;
@JsonIgnore
public List<ToolCallback> getToolCallbacks() {
return this.toolCallbacks;
}
public void setProxyToolCalls(Boolean proxyToolCalls) {
this.proxyToolCalls = proxyToolCalls;
@Override
@JsonIgnore
public void setToolCallbacks(List<ToolCallback> toolCallbacks) {
Assert.notNull(toolCallbacks, "toolCallbacks cannot be null");
Assert.noNullElements(toolCallbacks, "toolCallbacks cannot contain null elements");
this.toolCallbacks = toolCallbacks;
}
@Override
@JsonIgnore
public Set<String> getToolNames() {
return this.toolNames;
}
@Override
@JsonIgnore
public void setToolNames(Set<String> toolNames) {
Assert.notNull(toolNames, "toolNames cannot be null");
Assert.noNullElements(toolNames, "toolNames cannot contain null elements");
toolNames.forEach(tool -> Assert.hasText(tool, "toolNames cannot contain empty elements"));
this.toolNames = toolNames;
}
@Override
@Nullable
@JsonIgnore
public Boolean getInternalToolExecutionEnabled() {
return this.internalToolExecutionEnabled;
}
@Override
@JsonIgnore
public void setInternalToolExecutionEnabled(@Nullable Boolean internalToolExecutionEnabled) {
this.internalToolExecutionEnabled = internalToolExecutionEnabled;
}
@Override
@@ -357,7 +376,10 @@ public class MiniMaxChatOptions implements FunctionCallingOptions {
result = prime * result + ((this.maskSensitiveInfo == null) ? 0 : this.maskSensitiveInfo.hashCode());
result = prime * result + ((this.tools == null) ? 0 : this.tools.hashCode());
result = prime * result + ((this.toolChoice == null) ? 0 : this.toolChoice.hashCode());
result = prime * result + ((this.proxyToolCalls == null) ? 0 : this.proxyToolCalls.hashCode());
result = prime * result + ((this.toolCallbacks == null) ? 0 : this.toolCallbacks.hashCode());
result = prime * result + ((this.toolNames == null) ? 0 : this.toolNames.hashCode());
result = prime * result
+ ((this.internalToolExecutionEnabled == null) ? 0 : this.internalToolExecutionEnabled.hashCode());
result = prime * result + ((this.toolContext == null) ? 0 : this.toolContext.hashCode());
return result;
}
@@ -478,12 +500,30 @@ public class MiniMaxChatOptions implements FunctionCallingOptions {
else if (!this.toolChoice.equals(other.toolChoice)) {
return false;
}
if (this.proxyToolCalls == null) {
if (other.proxyToolCalls != null) {
if (this.internalToolExecutionEnabled == null) {
if (other.internalToolExecutionEnabled != null) {
return false;
}
}
else if (!this.proxyToolCalls.equals(other.proxyToolCalls)) {
else if (!this.internalToolExecutionEnabled.equals(other.internalToolExecutionEnabled)) {
return false;
}
if (this.toolNames == null) {
if (other.toolNames != null) {
return false;
}
}
else if (!this.toolNames.equals(other.toolNames)) {
return false;
}
if (this.toolCallbacks == null) {
if (other.toolCallbacks != null) {
return false;
}
}
else if (!this.toolCallbacks.equals(other.toolCallbacks)) {
return false;
}
@@ -581,25 +621,31 @@ public class MiniMaxChatOptions implements FunctionCallingOptions {
return this;
}
public Builder functionCallbacks(List<FunctionCallback> functionCallbacks) {
this.options.functionCallbacks = functionCallbacks;
public Builder toolCallbacks(List<ToolCallback> toolCallbacks) {
this.options.setToolCallbacks(toolCallbacks);
return this;
}
public Builder functions(Set<String> functionNames) {
Assert.notNull(functionNames, "Function names must not be null");
this.options.functions = functionNames;
public Builder toolCallbacks(ToolCallback... toolCallbacks) {
Assert.notNull(toolCallbacks, "toolCallbacks cannot be null");
this.options.toolCallbacks.addAll(Arrays.asList(toolCallbacks));
return this;
}
public Builder function(String functionName) {
Assert.hasText(functionName, "Function name must not be empty");
this.options.functions.add(functionName);
public Builder toolNames(Set<String> toolNames) {
Assert.notNull(toolNames, "toolNames cannot be null");
this.options.setToolNames(toolNames);
return this;
}
public Builder proxyToolCalls(Boolean proxyToolCalls) {
this.options.proxyToolCalls = proxyToolCalls;
public Builder toolNames(String... toolNames) {
Assert.notNull(toolNames, "toolNames cannot be null");
this.options.toolNames.addAll(Set.of(toolNames));
return this;
}
public Builder internalToolExecutionEnabled(@Nullable Boolean internalToolExecutionEnabled) {
this.options.setInternalToolExecutionEnabled(internalToolExecutionEnabled);
return this;
}

View File

@@ -23,13 +23,13 @@ import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.minimax.api.MiniMaxApi;
import org.springframework.ai.minimax.api.MockWeatherService;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.tool.function.FunctionToolCallback;
import static org.assertj.core.api.Assertions.assertThat;
/**
* @author Geng Rong
* @author Ilayaperumal Gopinathan
*/
public class ChatCompletionRequestTests {
@@ -39,7 +39,8 @@ public class ChatCompletionRequestTests {
var client = new MiniMaxChatModel(new MiniMaxApi("TEST"),
MiniMaxChatOptions.builder().model("DEFAULT_MODEL").temperature(66.6).build());
var request = client.createRequest(new Prompt("Test message content"), false);
var request = client.createRequest(new Prompt("Test message content",
MiniMaxChatOptions.builder().model("DEFAULT_MODEL").temperature(66.6).build()), false);
assertThat(request.messages()).hasSize(1);
assertThat(request.stream()).isFalse();
@@ -65,16 +66,15 @@ public class ChatCompletionRequestTests {
var client = new MiniMaxChatModel(new MiniMaxApi("TEST"),
MiniMaxChatOptions.builder().model("DEFAULT_MODEL").build());
var request = client.createRequest(new Prompt("Test message content", MiniMaxChatOptions.builder()
.model("PROMPT_MODEL")
.functionCallbacks(List.of(FunctionToolCallback.builder(TOOL_FUNCTION_NAME, new MockWeatherService())
.description("Get the weather in location")
.inputType(MockWeatherService.Request.class)
.build()))
.build()), false);
assertThat(client.getFunctionCallbackRegister()).hasSize(1);
assertThat(client.getFunctionCallbackRegister()).containsKeys(TOOL_FUNCTION_NAME);
var request = client.createRequest(new Prompt("Test message content",
MiniMaxChatOptions.builder()
.model("PROMPT_MODEL")
.toolCallbacks(List.of(FunctionToolCallback.builder(TOOL_FUNCTION_NAME, new MockWeatherService())
.description("Get the weather in location")
.inputType(MockWeatherService.Request.class)
.build()))
.build()),
false);
assertThat(request.messages()).hasSize(1);
assertThat(request.stream()).isFalse();
@@ -92,55 +92,19 @@ public class ChatCompletionRequestTests {
var client = new MiniMaxChatModel(new MiniMaxApi("TEST"),
MiniMaxChatOptions.builder()
.model("DEFAULT_MODEL")
.functionCallbacks(List.of(FunctionCallback.builder()
.function(TOOL_FUNCTION_NAME, new MockWeatherService())
.toolCallbacks(List.of(FunctionToolCallback.builder(TOOL_FUNCTION_NAME, new MockWeatherService())
.description("Get the weather in location")
.inputType(MockWeatherService.Request.class)
.build()))
.build());
var request = client.createRequest(new Prompt("Test message content"), false);
var prompt = client.buildRequestPrompt(new Prompt("Test message content"));
assertThat(client.getFunctionCallbackRegister()).hasSize(1);
assertThat(client.getFunctionCallbackRegister()).containsKeys(TOOL_FUNCTION_NAME);
assertThat(client.getFunctionCallbackRegister().get(TOOL_FUNCTION_NAME).getDescription())
.isEqualTo("Get the weather in location");
var request = client.createRequest(prompt, false);
assertThat(request.messages()).hasSize(1);
assertThat(request.stream()).isFalse();
assertThat(request.model()).isEqualTo("DEFAULT_MODEL");
assertThat(request.tools()).as("Default Options callback functions are not automatically enabled!")
.isNullOrEmpty();
// Explicitly enable the function
request = client.createRequest(
new Prompt("Test message content", MiniMaxChatOptions.builder().function(TOOL_FUNCTION_NAME).build()),
false);
assertThat(request.tools()).hasSize(1);
assertThat(request.tools().get(0).getFunction().getName()).as("Explicitly enabled function")
.isEqualTo(TOOL_FUNCTION_NAME);
// Override the default options function with one from the prompt
request = client.createRequest(new Prompt("Test message content",
MiniMaxChatOptions.builder()
.functionCallbacks(List.of(FunctionCallback.builder()
.function(TOOL_FUNCTION_NAME, new MockWeatherService())
.description("Overridden function description")
.inputType(MockWeatherService.Request.class)
.build()))
.build()),
false);
assertThat(request.tools()).hasSize(1);
assertThat(request.tools().get(0).getFunction().getName()).as("Explicitly enabled function")
.isEqualTo(TOOL_FUNCTION_NAME);
assertThat(client.getFunctionCallbackRegister()).hasSize(1);
assertThat(client.getFunctionCallbackRegister()).containsKeys(TOOL_FUNCTION_NAME);
assertThat(client.getFunctionCallbackRegister().get(TOOL_FUNCTION_NAME).getDescription())
.isEqualTo("Overridden function description");
}
}

View File

@@ -26,6 +26,7 @@ import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;
import reactor.core.publisher.Flux;
import org.springframework.ai.chat.prompt.ChatOptions;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.document.MetadataMode;
import org.springframework.ai.minimax.MiniMaxChatModel;
@@ -40,6 +41,7 @@ import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletionMessage.Role;
import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletionRequest;
import org.springframework.ai.minimax.api.MiniMaxApi.EmbeddingList;
import org.springframework.ai.minimax.api.MiniMaxApi.EmbeddingRequest;
import org.springframework.ai.model.tool.ToolCallingManager;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.ai.retry.TransientAiException;
import org.springframework.http.ResponseEntity;
@@ -76,8 +78,8 @@ public class MiniMaxRetryTests {
this.retryListener = new TestRetryListener();
this.retryTemplate.registerListener(this.retryListener);
this.chatModel = new MiniMaxChatModel(this.miniMaxApi, MiniMaxChatOptions.builder().build(), null,
this.retryTemplate);
this.chatModel = new MiniMaxChatModel(this.miniMaxApi, MiniMaxChatOptions.builder().build(),
ToolCallingManager.builder().build(), this.retryTemplate);
this.embeddingModel = new MiniMaxEmbeddingModel(this.miniMaxApi, MetadataMode.EMBED,
MiniMaxEmbeddingOptions.builder().build(), this.retryTemplate);
}
@@ -95,7 +97,7 @@ public class MiniMaxRetryTests {
.willThrow(new TransientAiException("Transient Error 2"))
.willReturn(ResponseEntity.of(Optional.of(expectedChatCompletion)));
var result = this.chatModel.call(new Prompt("text"));
var result = this.chatModel.call(new Prompt("text", ChatOptions.builder().build()));
assertThat(result).isNotNull();
assertThat(result.getResult().getOutput().getText()).isSameAs("Response");
@@ -123,7 +125,7 @@ public class MiniMaxRetryTests {
.willThrow(new TransientAiException("Transient Error 2"))
.willReturn(Flux.just(expectedChatCompletion));
var result = this.chatModel.stream(new Prompt("text"));
var result = this.chatModel.stream(new Prompt("text", ChatOptions.builder().build()));
assertThat(result).isNotNull();
assertThat(result.collectList().block().get(0).getResult().getOutput().getText()).isSameAs("Response");

View File

@@ -33,14 +33,15 @@ import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.minimax.MiniMaxChatModel;
import org.springframework.ai.minimax.MiniMaxChatOptions;
import org.springframework.ai.minimax.api.MiniMaxApi;
import org.springframework.ai.model.function.DefaultFunctionCallbackResolver;
import org.springframework.ai.model.tool.DefaultToolExecutionEligibilityPredicate;
import org.springframework.ai.model.tool.ToolCallingManager;
import org.springframework.ai.observation.conventions.AiOperationType;
import org.springframework.ai.observation.conventions.AiProvider;
import org.springframework.ai.retry.RetryUtils;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.context.annotation.Bean;
import org.springframework.retry.support.RetryTemplate;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.chat.observation.ChatModelObservationDocumentation.HighCardinalityKeyNames;
@@ -171,8 +172,8 @@ public class MiniMaxChatModelObservationIT {
@Bean
public MiniMaxChatModel minimaxChatModel(MiniMaxApi minimaxApi, TestObservationRegistry observationRegistry) {
return new MiniMaxChatModel(minimaxApi, MiniMaxChatOptions.builder().build(),
new DefaultFunctionCallbackResolver(), List.of(), RetryTemplate.defaultInstance(),
observationRegistry);
ToolCallingManager.builder().build(), RetryUtils.DEFAULT_RETRY_TEMPLATE, observationRegistry,
new DefaultToolExecutionEligibilityPredicate());
}
}

View File

@@ -36,11 +36,14 @@ import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.minimax.MiniMaxChatModel;
import org.springframework.ai.minimax.MiniMaxChatOptions;
import org.springframework.ai.minimax.api.MiniMaxApi;
import org.springframework.ai.minimax.api.MockWeatherService;
import org.springframework.ai.tool.function.FunctionToolCallback;
import static org.assertj.core.api.Assertions.assertThat;
/**
* @author Geng Rong
* @author Ilayaperumal Gopinathan
*/
@EnabledIfEnvironmentVariable(named = "MINIMAX_API_KEY", matches = ".+")
public class MiniMaxChatOptionsTests {
@@ -58,72 +61,45 @@ public class MiniMaxChatOptionsTests {
List<Message> messages = new ArrayList<>(List.of(userMessage));
// markSensitiveInfo is enabled by default
ChatResponse response = this.chatModel.call(new Prompt(messages));
ChatResponse response = this.chatModel
.call(new Prompt(messages, MiniMaxChatOptions.builder().maskSensitiveInfo(true).build()));
String responseContent = response.getResult().getOutput().getText();
assertThat(responseContent).contains("133-**");
assertThat(responseContent).doesNotContain("133-12345678");
var chatOptions = MiniMaxChatOptions.builder().maskSensitiveInfo(false).build();
ChatResponse unmaskResponse = this.chatModel.call(new Prompt(messages, chatOptions));
String unmaskResponseContent = unmaskResponse.getResult().getOutput().getText();
assertThat(unmaskResponseContent).contains("133-12345678");
}
/**
* There is a certain probability of failure, because it needs to be searched through
* the network, which may cause the test to fail due to different search results. And
* the search results are related to time. For example, after the start of the Paris
* Paralympic Games, searching for the number of gold medals in the Paris Olympics may
* be affected by the search results of the number of gold medals in the Paris
* Paralympic Games with higher priority by the search engine. Even if the input is an
* English question, there may be get Chinese content, because the main training
* content of MiniMax and search engine are Chinese
*/
@Test
void testWebSearch() {
UserMessage userMessage = new UserMessage(
"How many gold medals has the United States won in total at the 2024 Olympics?");
void testToolCalling() {
UserMessage userMessage = new UserMessage("What is the weather in San Francisco?");
List<Message> messages = new ArrayList<>(List.of(userMessage));
List<MiniMaxApi.FunctionTool> functionTool = List.of(MiniMaxApi.FunctionTool.webSearchFunctionTool());
MiniMaxChatOptions options = MiniMaxChatOptions.builder()
.model(org.springframework.ai.minimax.api.MiniMaxApi.ChatModel.ABAB_6_5_S_Chat.value)
.tools(functionTool)
.toolCallbacks(List.of(FunctionToolCallback.builder("CurrentWeather", new MockWeatherService())
.description("Get the weather in location")
.inputType(MockWeatherService.Request.class)
.build()))
.build();
ChatResponse response = this.chatModel.call(new Prompt(messages, options));
String responseContent = response.getResult().getOutput().getText();
assertThat(responseContent).contains("40");
assertThat(responseContent).contains("30");
}
/**
* There is a certain probability of failure, because it needs to be searched through
* the network, which may cause the test to fail due to different search results. And
* the search results are related to time. For example, after the start of the Paris
* Paralympic Games, searching for the number of gold medals in the Paris Olympics may
* be affected by the search results of the number of gold medals in the Paris
* Paralympic Games with higher priority by the search engine. Even if the input is an
* English question, there may be get Chinese content, because the main training
* content of MiniMax and search engine of MiniMax are Chinese
*/
@Test
void testWebSearchStream() {
UserMessage userMessage = new UserMessage(
"How many gold medals has the United States won in total at the 2024 Olympics?");
void testToolCallingStream() {
UserMessage userMessage = new UserMessage("What is the weather in Paris?");
List<Message> messages = new ArrayList<>(List.of(userMessage));
List<MiniMaxApi.FunctionTool> functionTool = List.of(MiniMaxApi.FunctionTool.webSearchFunctionTool());
MiniMaxChatOptions options = MiniMaxChatOptions.builder()
.model(org.springframework.ai.minimax.api.MiniMaxApi.ChatModel.ABAB_6_5_S_Chat.value)
.tools(functionTool)
.toolCallbacks(List.of(FunctionToolCallback.builder("CurrentWeather", new MockWeatherService())
.description("Get the weather in location")
.inputType(MockWeatherService.Request.class)
.build()))
.build();
Flux<ChatResponse> response = this.chatModel.stream(new Prompt(messages, options));
@@ -137,7 +113,7 @@ public class MiniMaxChatOptionsTests {
.collect(Collectors.joining());
logger.info("Response: {}", content);
assertThat(content).contains("40");
assertThat(content).contains("15");
}
}