Streamline the VertexAI Gemini Function Calling
- Align with AbstractToolCallSupport
This commit is contained in:
@@ -15,9 +15,43 @@
|
||||
*/
|
||||
package org.springframework.ai.vertexai.gemini;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collection;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.messages.MessageType;
|
||||
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.metadata.ChatResponseMetadata;
|
||||
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;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.ChatModelDescription;
|
||||
import org.springframework.ai.model.Media;
|
||||
import org.springframework.ai.model.ModelOptionsUtils;
|
||||
import org.springframework.ai.model.function.FunctionCallbackContext;
|
||||
import org.springframework.ai.vertexai.gemini.metadata.VertexAiUsage;
|
||||
import org.springframework.beans.factory.DisposableBean;
|
||||
import org.springframework.lang.NonNull;
|
||||
import org.springframework.util.Assert;
|
||||
import org.springframework.util.CollectionUtils;
|
||||
import org.springframework.util.StringUtils;
|
||||
|
||||
import com.fasterxml.jackson.annotation.JsonInclude;
|
||||
import com.fasterxml.jackson.annotation.JsonInclude.Include;
|
||||
import com.google.cloud.vertexai.VertexAI;
|
||||
import com.google.cloud.vertexai.api.Candidate;
|
||||
import com.google.cloud.vertexai.api.Candidate.FinishReason;
|
||||
import com.google.cloud.vertexai.api.Content;
|
||||
import com.google.cloud.vertexai.api.FunctionCall;
|
||||
import com.google.cloud.vertexai.api.FunctionDeclaration;
|
||||
@@ -32,38 +66,8 @@ import com.google.cloud.vertexai.generativeai.PartMaker;
|
||||
import com.google.cloud.vertexai.generativeai.ResponseStream;
|
||||
import com.google.protobuf.Struct;
|
||||
import com.google.protobuf.util.JsonFormat;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.model.Media;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.messages.MessageType;
|
||||
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.ChatResponseMetadata;
|
||||
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.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.ChatModelDescription;
|
||||
import org.springframework.ai.model.ModelOptionsUtils;
|
||||
import org.springframework.ai.chat.model.AbstractToolCallSupport;
|
||||
import org.springframework.ai.model.function.FunctionCallbackContext;
|
||||
import org.springframework.ai.vertexai.gemini.metadata.VertexAiUsage;
|
||||
import org.springframework.beans.factory.DisposableBean;
|
||||
import org.springframework.lang.NonNull;
|
||||
import org.springframework.util.Assert;
|
||||
import org.springframework.util.CollectionUtils;
|
||||
import org.springframework.util.StringUtils;
|
||||
import reactor.core.publisher.Flux;
|
||||
import reactor.core.publisher.Mono;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collection;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
/**
|
||||
* @author Christian Tzolov
|
||||
@@ -161,47 +165,22 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
|
||||
|
||||
GenerateContentResponse response = this.getContentResponse(geminiRequest);
|
||||
|
||||
if (this.isToolFunctionCall(response)) {
|
||||
List<Message> toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(), response);
|
||||
return this.call(new Prompt(toolCallMessageConversation, prompt.getOptions()));
|
||||
|
||||
}
|
||||
|
||||
List<Generation> generations = response.getCandidatesList()
|
||||
.stream()
|
||||
.map(candidate -> candidate.getContent().getPartsList())
|
||||
.map(this::responseCandiateToGeneration)
|
||||
.flatMap(List::stream)
|
||||
.map(Part::getText)
|
||||
.map(t -> new Generation(t))
|
||||
.toList();
|
||||
|
||||
return new ChatResponse(generations, toChatResponseMetadata(response));
|
||||
}
|
||||
ChatResponse chatResponse = new ChatResponse(generations, toChatResponseMetadata(response));
|
||||
|
||||
public List<Message> handleToolCallRequests(List<Message> previousMessages, GenerateContentResponse response) {
|
||||
if (isToolCall(chatResponse, Set.of(FinishReason.STOP.name()))) {
|
||||
var toolCallConversation = handleToolCalls(prompt, chatResponse);
|
||||
// Recursively call the call method with the tool call message
|
||||
// conversation that contains the call responses.
|
||||
return this.call(new Prompt(toolCallConversation, prompt.getOptions()));
|
||||
}
|
||||
|
||||
Content assistantContent = response.getCandidatesList().get(0).getContent();
|
||||
|
||||
List<AssistantMessage.ToolCall> assistantToolCalls = assistantContent.getPartsList()
|
||||
.stream()
|
||||
.filter(part -> part.hasFunctionCall())
|
||||
.map(part -> {
|
||||
FunctionCall functionCall = part.getFunctionCall();
|
||||
var functionName = functionCall.getName();
|
||||
String functionArguments = structToJson(functionCall.getArgs());
|
||||
return new AssistantMessage.ToolCall("", "function", functionName, functionArguments);
|
||||
})
|
||||
.toList();
|
||||
|
||||
AssistantMessage assistantMessage = new AssistantMessage("", Map.of(), assistantToolCalls);
|
||||
|
||||
ToolResponseMessage toolResponseMessage = this.executeFunctions(assistantMessage);
|
||||
|
||||
// History
|
||||
List<Message> toolCallMessageConversation = new ArrayList<>(previousMessages);
|
||||
toolCallMessageConversation.add(assistantMessage);
|
||||
toolCallMessageConversation.add(toolResponseMessage);
|
||||
return toolCallMessageConversation;
|
||||
return chatResponse;
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -214,26 +193,24 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
|
||||
.generateContentStream(request.contents);
|
||||
|
||||
return Flux.fromStream(responseStream.stream()).switchMap(response -> {
|
||||
if (this.isToolFunctionCall(response)) {
|
||||
List<Message> toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(),
|
||||
response);
|
||||
|
||||
List<Generation> generations = response.getCandidatesList()
|
||||
.stream()
|
||||
.map(this::responseCandiateToGeneration)
|
||||
.flatMap(List::stream)
|
||||
.toList();
|
||||
|
||||
ChatResponse chatResponse = new ChatResponse(generations, toChatResponseMetadata(response));
|
||||
|
||||
if (isToolCall(chatResponse,
|
||||
Set.of(FinishReason.STOP.name(), FinishReason.FINISH_REASON_UNSPECIFIED.name()))) {
|
||||
var toolCallConversation = handleToolCalls(prompt, chatResponse);
|
||||
// Recursively call the stream method with the tool call message
|
||||
// conversation that contains the call responses.
|
||||
return this.stream(new Prompt(toolCallMessageConversation, prompt.getOptions()));
|
||||
return this.stream(new Prompt(toolCallConversation, prompt.getOptions()));
|
||||
}
|
||||
|
||||
return Mono.just(response).map(response2 -> {
|
||||
List<Generation> generations = response.getCandidatesList()
|
||||
.stream()
|
||||
.map(candidate -> candidate.getContent().getPartsList())
|
||||
.flatMap(List::stream)
|
||||
.map(Part::getText)
|
||||
.map(t -> new Generation(t))
|
||||
.toList();
|
||||
|
||||
return new ChatResponse(generations, toChatResponseMetadata(response));
|
||||
|
||||
});
|
||||
return Flux.just(chatResponse);
|
||||
});
|
||||
}
|
||||
catch (Exception e) {
|
||||
@@ -241,6 +218,49 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
|
||||
}
|
||||
}
|
||||
|
||||
protected List<Generation> responseCandiateToGeneration(Candidate candidate) {
|
||||
|
||||
// TODO - The candidateIndex (e.g. choice must be asigned to the generation).
|
||||
int candidateIndex = candidate.getIndex();
|
||||
FinishReason candidateFinishReasonn = candidate.getFinishReason();
|
||||
|
||||
Map<String, Object> messageMetadata = Map.of("candidateIndex", candidateIndex, "finishReason",
|
||||
candidateFinishReasonn);
|
||||
|
||||
ChatGenerationMetadata chatGenerationMetadata = ChatGenerationMetadata.from(candidateFinishReasonn.name(),
|
||||
null);
|
||||
|
||||
boolean isFunctinCall = candidate.getContent().getPartsList().stream().allMatch(Part::hasFunctionCall);
|
||||
|
||||
if (isFunctinCall) {
|
||||
List<AssistantMessage.ToolCall> assistantToolCalls = candidate.getContent()
|
||||
.getPartsList()
|
||||
.stream()
|
||||
.filter(part -> part.hasFunctionCall())
|
||||
.map(part -> {
|
||||
FunctionCall functionCall = part.getFunctionCall();
|
||||
var functionName = functionCall.getName();
|
||||
String functionArguments = structToJson(functionCall.getArgs());
|
||||
return new AssistantMessage.ToolCall("", "function", functionName, functionArguments);
|
||||
})
|
||||
.toList();
|
||||
|
||||
AssistantMessage assistantMessage = new AssistantMessage("", messageMetadata, assistantToolCalls);
|
||||
|
||||
return List.of(new Generation(assistantMessage, chatGenerationMetadata));
|
||||
}
|
||||
else {
|
||||
List<Generation> generations = candidate.getContent()
|
||||
.getPartsList()
|
||||
.stream()
|
||||
.map(part -> new AssistantMessage(part.getText(), messageMetadata))
|
||||
.map(assistantMessage -> new Generation(assistantMessage, chatGenerationMetadata))
|
||||
.toList();
|
||||
|
||||
return generations;
|
||||
}
|
||||
}
|
||||
|
||||
private ChatResponseMetadata toChatResponseMetadata(GenerateContentResponse response) {
|
||||
return ChatResponseMetadata.builder().withUsage(new VertexAiUsage(response.getUsageMetadata())).build();
|
||||
}
|
||||
@@ -499,15 +519,6 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements
|
||||
}
|
||||
}
|
||||
|
||||
protected boolean isToolFunctionCall(GenerateContentResponse response) {
|
||||
if (response == null || CollectionUtils.isEmpty(response.getCandidatesList())
|
||||
|| response.getCandidatesList().get(0).getContent() == null
|
||||
|| CollectionUtils.isEmpty(response.getCandidatesList().get(0).getContent().getPartsList())) {
|
||||
return false;
|
||||
}
|
||||
return response.getCandidatesList().get(0).getContent().getPartsList().get(0).hasFunctionCall();
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatOptions getDefaultOptions() {
|
||||
return VertexAiGeminiChatOptions.fromOptions(this.defaultOptions);
|
||||
|
||||
@@ -29,7 +29,10 @@ import java.util.concurrent.ConcurrentHashMap;
|
||||
/**
|
||||
* @author Christian Tzolov
|
||||
* @author Grogdunn
|
||||
* @deprecated since 1.0.0-M1 in favor of
|
||||
* {@link org.springframework.ai.chat.model.AbstractToolCallSupport}
|
||||
*/
|
||||
@Deprecated(since = "1.0.0-M2", forRemoval = true)
|
||||
public abstract class AbstractFunctionCallSupport<Msg, Req, Resp> {
|
||||
|
||||
protected final static boolean IS_RUNTIME_CALL = true;
|
||||
|
||||
Reference in New Issue
Block a user