Add high-level function calling support for Mistral AI
This commit is contained in:
@@ -15,12 +15,25 @@
|
||||
*/
|
||||
package org.springframework.ai.mistralai;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.HashMap;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
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.SystemMessage;
|
||||
import org.springframework.ai.chat.messages.ToolResponseMessage;
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
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.metadata.ChatGenerationMetadata;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.mistralai.api.MistralAiApi;
|
||||
@@ -28,26 +41,21 @@ import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletion;
|
||||
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletion.Choice;
|
||||
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionChunk;
|
||||
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionMessage;
|
||||
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionMessage.ChatCompletionFunction;
|
||||
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionMessage.ToolCall;
|
||||
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionRequest;
|
||||
import org.springframework.ai.mistralai.metadata.MistralAiChatResponseMetadata;
|
||||
import org.springframework.ai.model.ModelOptionsUtils;
|
||||
import org.springframework.ai.model.function.AbstractFunctionCallSupport;
|
||||
import org.springframework.ai.model.function.AbstractToolCallSupport;
|
||||
import org.springframework.ai.model.function.FunctionCallbackContext;
|
||||
import org.springframework.ai.retry.RetryUtils;
|
||||
import org.springframework.http.ResponseEntity;
|
||||
import org.springframework.retry.support.RetryTemplate;
|
||||
import org.springframework.util.Assert;
|
||||
import org.springframework.util.CollectionUtils;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
import java.util.HashMap;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Optional;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import reactor.core.publisher.Flux;
|
||||
import reactor.core.publisher.Mono;
|
||||
|
||||
/**
|
||||
* @author Ricken Bazolo
|
||||
@@ -57,9 +65,7 @@ import java.util.concurrent.ConcurrentHashMap;
|
||||
* @author luocongqiu
|
||||
* @since 0.8.1
|
||||
*/
|
||||
public class MistralAiChatModel extends
|
||||
AbstractFunctionCallSupport<MistralAiApi.ChatCompletionMessage, MistralAiApi.ChatCompletionRequest, ResponseEntity<MistralAiApi.ChatCompletion>>
|
||||
implements ChatModel {
|
||||
public class MistralAiChatModel extends AbstractToolCallSupport<ChatCompletion> implements ChatModel {
|
||||
|
||||
private final Logger log = LoggerFactory.getLogger(getClass());
|
||||
|
||||
@@ -106,7 +112,15 @@ public class MistralAiChatModel extends
|
||||
|
||||
return retryTemplate.execute(ctx -> {
|
||||
|
||||
ResponseEntity<ChatCompletion> completionEntity = this.callWithFunctionSupport(request);
|
||||
ResponseEntity<ChatCompletion> completionEntity = this.mistralAiApi.chatCompletionEntity(request);
|
||||
|
||||
if (this.isToolFunctionCall(completionEntity.getBody())) {
|
||||
List<Message> toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(),
|
||||
completionEntity.getBody());
|
||||
// Recursively call the call method with the tool call message
|
||||
// conversation that contains the call responses.
|
||||
return this.call(new Prompt(toolCallMessageConversation, prompt.getOptions()));
|
||||
}
|
||||
|
||||
var chatCompletion = completionEntity.getBody();
|
||||
if (chatCompletion == null) {
|
||||
@@ -144,21 +158,26 @@ public class MistralAiChatModel extends
|
||||
|
||||
return retryTemplate.execute(ctx -> {
|
||||
|
||||
var completionChunks = this.mistralAiApi.chatCompletionStream(request);
|
||||
Flux<ChatCompletionChunk> completionChunks = this.mistralAiApi.chatCompletionStream(request);
|
||||
|
||||
// For chunked responses, only the first chunk contains the choice role.
|
||||
// The rest of the chunks with same ID share the same role.
|
||||
ConcurrentHashMap<String, String> roleMap = new ConcurrentHashMap<>();
|
||||
|
||||
return completionChunks.map(chunk -> toChatCompletion(chunk))
|
||||
.switchMap(
|
||||
cc -> handleFunctionCallOrReturnStream(request, Flux.just(ResponseEntity.of(Optional.of(cc)))))
|
||||
.map(ResponseEntity::getBody)
|
||||
.map(chatCompletion -> {
|
||||
@SuppressWarnings("null")
|
||||
String id = chatCompletion.id();
|
||||
return completionChunks.map(this::toChatCompletion).switchMap(chatCompletion -> {
|
||||
if (this.isToolFunctionCall(chatCompletion)) {
|
||||
var toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(),
|
||||
chatCompletion);
|
||||
// Recursively call the stream method with the tool call message
|
||||
// conversation that contains the call responses.
|
||||
return this.stream(new Prompt(toolCallMessageConversation, prompt.getOptions()));
|
||||
}
|
||||
|
||||
List<Generation> generations = chatCompletion.choices().stream().map(choice -> {
|
||||
return Mono.just(chatCompletion).map(chatCompletion2 -> {
|
||||
@SuppressWarnings("null")
|
||||
String id = chatCompletion2.id();
|
||||
|
||||
List<Generation> generations = chatCompletion2.choices().stream().map(choice -> {
|
||||
if (choice.message().role() != null) {
|
||||
roleMap.putIfAbsent(id, choice.message().role().name());
|
||||
}
|
||||
@@ -171,11 +190,67 @@ public class MistralAiChatModel extends
|
||||
}
|
||||
return generation;
|
||||
}).toList();
|
||||
|
||||
return new ChatResponse(generations);
|
||||
});
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
private List<Message> handleToolCallRequests(List<Message> previousMessages, ChatCompletion chatCompletion) {
|
||||
|
||||
ChatCompletionMessage nativeAssistantMessage = this.extractAssistantMessage(chatCompletion);
|
||||
|
||||
List<AssistantMessage.ToolCall> assistantToolCalls = nativeAssistantMessage.toolCalls()
|
||||
.stream()
|
||||
.map(toolCall -> new AssistantMessage.ToolCall(toolCall.id(), "function", toolCall.function().name(),
|
||||
toolCall.function().arguments()))
|
||||
.toList();
|
||||
|
||||
AssistantMessage assistantMessage = new AssistantMessage(nativeAssistantMessage.content(), Map.of(),
|
||||
assistantToolCalls);
|
||||
|
||||
ToolResponseMessage toolResponseMessage = this.executeFuncitons(assistantMessage);
|
||||
|
||||
// History
|
||||
List<Message> messages = new ArrayList<>(previousMessages);
|
||||
messages.add(assistantMessage);
|
||||
messages.add(toolResponseMessage);
|
||||
|
||||
return messages;
|
||||
}
|
||||
|
||||
private ChatCompletionMessage extractAssistantMessage(ChatCompletion chatCompletion) {
|
||||
ChatCompletionMessage msg = chatCompletion.choices().iterator().next().message();
|
||||
if (msg.role() == null) {
|
||||
// add missing role
|
||||
msg = new ChatCompletionMessage(msg.content(), ChatCompletionMessage.Role.ASSISTANT, msg.name(),
|
||||
msg.toolCalls());
|
||||
}
|
||||
return msg;
|
||||
}
|
||||
|
||||
protected ToolResponseMessage executeFuncitons(AssistantMessage assistantMessage) {
|
||||
|
||||
List<ToolResponseMessage.ToolResponse> toolResponses = new ArrayList<>();
|
||||
|
||||
for (AssistantMessage.ToolCall toolCall : assistantMessage.getToolCalls()) {
|
||||
|
||||
var functionName = toolCall.name();
|
||||
String functionArguments = toolCall.arguments();
|
||||
|
||||
if (!this.functionCallbackRegister.containsKey(functionName)) {
|
||||
throw new IllegalStateException("No function callback found for function name: " + functionName);
|
||||
}
|
||||
|
||||
String functionResponse = this.functionCallbackRegister.get(functionName).call(functionArguments);
|
||||
|
||||
toolResponses.add(new ToolResponseMessage.ToolResponse(toolCall.id(), functionName, functionResponse));
|
||||
}
|
||||
|
||||
return new ToolResponseMessage(toolResponses, Map.of());
|
||||
}
|
||||
|
||||
private ChatCompletion toChatCompletion(ChatCompletionChunk chunk) {
|
||||
List<Choice> choices = chunk.choices()
|
||||
.stream()
|
||||
@@ -192,11 +267,44 @@ public class MistralAiChatModel extends
|
||||
|
||||
Set<String> functionsForThisRequest = new HashSet<>();
|
||||
|
||||
var chatCompletionMessages = prompt.getInstructions()
|
||||
.stream()
|
||||
.map(m -> new MistralAiApi.ChatCompletionMessage(m.getContent(),
|
||||
MistralAiApi.ChatCompletionMessage.Role.valueOf(m.getMessageType().name())))
|
||||
.toList();
|
||||
List<ChatCompletionMessage> chatCompletionMessages = prompt.getInstructions().stream().map(message -> {
|
||||
if (message instanceof UserMessage userMessage) {
|
||||
return List.of(new MistralAiApi.ChatCompletionMessage(userMessage.getContent(),
|
||||
MistralAiApi.ChatCompletionMessage.Role.USER));
|
||||
}
|
||||
else if (message instanceof SystemMessage systemMessage) {
|
||||
return List.of(new MistralAiApi.ChatCompletionMessage(systemMessage.getContent(),
|
||||
MistralAiApi.ChatCompletionMessage.Role.SYSTEM));
|
||||
}
|
||||
else if (message instanceof AssistantMessage assistantMessage) {
|
||||
List<ToolCall> toolCalls = null;
|
||||
if (!CollectionUtils.isEmpty(assistantMessage.getToolCalls())) {
|
||||
toolCalls = assistantMessage.getToolCalls().stream().map(toolCall -> {
|
||||
var function = new ChatCompletionFunction(toolCall.name(), toolCall.arguments());
|
||||
return new ToolCall(toolCall.id(), toolCall.type(), function);
|
||||
}).toList();
|
||||
}
|
||||
|
||||
return List.of(new MistralAiApi.ChatCompletionMessage(assistantMessage.getContent(),
|
||||
MistralAiApi.ChatCompletionMessage.Role.ASSISTANT, null, toolCalls, null));
|
||||
}
|
||||
else if (message instanceof ToolResponseMessage toolResponseMessage) {
|
||||
|
||||
toolResponseMessage.getResponses().forEach(response -> {
|
||||
Assert.isTrue(response.id() != null, "ToolResponseMessage must have an id");
|
||||
Assert.isTrue(response.name() != null, "ToolResponseMessage must have a name");
|
||||
});
|
||||
|
||||
return toolResponseMessage.getResponses()
|
||||
.stream()
|
||||
.map(toolResponse -> new MistralAiApi.ChatCompletionMessage(toolResponse.responseData(),
|
||||
MistralAiApi.ChatCompletionMessage.Role.TOOL, toolResponse.name(), null, toolResponse.id()))
|
||||
.toList();
|
||||
}
|
||||
else {
|
||||
throw new IllegalStateException("Unexpected message type: " + message);
|
||||
}
|
||||
}).flatMap(List::stream).toList();
|
||||
|
||||
var request = new MistralAiApi.ChatCompletionRequest(chatCompletionMessages, stream);
|
||||
|
||||
@@ -239,74 +347,10 @@ public class MistralAiChatModel extends
|
||||
}).toList();
|
||||
}
|
||||
|
||||
//
|
||||
// Function Calling Support
|
||||
//
|
||||
@Override
|
||||
protected ChatCompletionRequest doCreateToolResponseRequest(ChatCompletionRequest previousRequest,
|
||||
ChatCompletionMessage responseMessage, List<ChatCompletionMessage> conversationHistory) {
|
||||
protected boolean isToolFunctionCall(ChatCompletion chatCompletion) {
|
||||
|
||||
// Every tool-call item requires a separate function call and a response (TOOL)
|
||||
// message.
|
||||
for (ToolCall toolCall : responseMessage.toolCalls()) {
|
||||
|
||||
String id = toolCall.id();
|
||||
String functionName = toolCall.function().name();
|
||||
String functionArguments = toolCall.function().arguments();
|
||||
|
||||
if (!this.functionCallbackRegister.containsKey(functionName)) {
|
||||
throw new IllegalStateException("No function callback found for function name: " + functionName);
|
||||
}
|
||||
|
||||
String functionResponse = this.functionCallbackRegister.get(functionName).call(functionArguments);
|
||||
|
||||
// Add the function response to the conversation.
|
||||
conversationHistory.add(new ChatCompletionMessage(functionResponse, ChatCompletionMessage.Role.TOOL,
|
||||
functionName, null, id));
|
||||
}
|
||||
|
||||
// Recursively call chatCompletionWithTools until the model doesn't call a
|
||||
// functions anymore.
|
||||
ChatCompletionRequest newRequest = new ChatCompletionRequest(conversationHistory, previousRequest.stream());
|
||||
newRequest = ModelOptionsUtils.merge(newRequest, previousRequest, ChatCompletionRequest.class);
|
||||
|
||||
return newRequest;
|
||||
}
|
||||
|
||||
@Override
|
||||
protected List<ChatCompletionMessage> doGetUserMessages(ChatCompletionRequest request) {
|
||||
return request.messages();
|
||||
}
|
||||
|
||||
@SuppressWarnings("null")
|
||||
@Override
|
||||
protected ChatCompletionMessage doGetToolResponseMessage(ResponseEntity<ChatCompletion> chatCompletion) {
|
||||
ChatCompletionMessage msg = chatCompletion.getBody().choices().iterator().next().message();
|
||||
if (msg.role() == null) {
|
||||
// add missing role
|
||||
msg = new ChatCompletionMessage(msg.content(), ChatCompletionMessage.Role.ASSISTANT, msg.name(),
|
||||
msg.toolCalls());
|
||||
}
|
||||
return msg;
|
||||
}
|
||||
|
||||
@Override
|
||||
protected ResponseEntity<ChatCompletion> doChatCompletion(ChatCompletionRequest request) {
|
||||
return this.mistralAiApi.chatCompletionEntity(request);
|
||||
}
|
||||
|
||||
@Override
|
||||
protected Flux<ResponseEntity<ChatCompletion>> doChatCompletionStream(ChatCompletionRequest request) {
|
||||
return this.mistralAiApi.chatCompletionStream(request)
|
||||
.map(this::toChatCompletion)
|
||||
.map(Optional::ofNullable)
|
||||
.map(ResponseEntity::of);
|
||||
}
|
||||
|
||||
@Override
|
||||
protected boolean isToolFunctionCall(ResponseEntity<ChatCompletion> chatCompletion) {
|
||||
|
||||
var body = chatCompletion.getBody();
|
||||
var body = chatCompletion;
|
||||
if (body == null) {
|
||||
return false;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user