Add high-level function calling support for Mistral AI

This commit is contained in:
Christian Tzolov
2024-07-16 16:18:16 +02:00
parent 7c26c7bcdf
commit 18cdeee092

View File

@@ -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;
}