Add high-level function calling support for Anthropic

- simplify and unify the AbstractToolCallSupport::executeFuncitons. Fix filed typos
 - remove old classes
This commit is contained in:
Christian Tzolov
2024-07-14 11:57:15 +02:00
parent 7e98fb78b7
commit ee19ee1db9
6 changed files with 133 additions and 612 deletions

View File

@@ -20,7 +20,6 @@ import java.util.Base64;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.Set;
import java.util.stream.Collectors;
@@ -34,7 +33,10 @@ import org.springframework.ai.anthropic.api.AnthropicApi.ContentBlock;
import org.springframework.ai.anthropic.api.AnthropicApi.ContentBlock.ContentBlockType;
import org.springframework.ai.anthropic.api.AnthropicApi.Role;
import org.springframework.ai.anthropic.metadata.AnthropicChatResponseMetadata;
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.ToolResponseMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
@@ -42,15 +44,17 @@ 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.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 org.springframework.util.StringUtils;
import reactor.core.publisher.Flux;
import reactor.core.publisher.Mono;
/**
* The {@link ChatModel} implementation for the Anthropic service.
@@ -60,13 +64,11 @@ import reactor.core.publisher.Flux;
* @author Mariusz Bernacki
* @since 1.0.0
*/
public class AnthropicChatModel extends
AbstractFunctionCallSupport<AnthropicApi.AnthropicMessage, AnthropicApi.ChatCompletionRequest, ResponseEntity<AnthropicApi.ChatCompletionResponse>>
implements ChatModel {
public class AnthropicChatModel extends AbstractToolCallSupport<ChatCompletionResponse> implements ChatModel {
private static final Logger logger = LoggerFactory.getLogger(AnthropicChatModel.class);
public static final String DEFAULT_MODEL_NAME = AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue();
public static final String DEFAULT_MODEL_NAME = AnthropicApi.ChatModel.CLAUDE_3_5_SONNET.getValue();
public static final Integer DEFAULT_MAX_TOKENS = 500;
@@ -148,7 +150,14 @@ public class AnthropicChatModel extends
ChatCompletionRequest request = createRequest(prompt, false);
return this.retryTemplate.execute(ctx -> {
ResponseEntity<ChatCompletionResponse> completionEntity = this.callWithFunctionSupport(request);
ResponseEntity<ChatCompletionResponse> completionEntity = this.anthropicApi.chatCompletionEntity(request);
if (this.isToolFunctionCall(completionEntity.getBody())) {
List<Message> toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(),
completionEntity.getBody());
return this.call(new Prompt(toolCallMessageConversation, prompt.getOptions()));
}
return toChatResponse(completionEntity.getBody());
});
}
@@ -162,14 +171,52 @@ public class AnthropicChatModel extends
Flux<ChatCompletionResponse> response = this.anthropicApi.chatCompletionStream(request);
return response
.switchMap(chatCompletionResponse -> handleFunctionCallOrReturnStream(request,
Flux.just(ResponseEntity.of(Optional.of(chatCompletionResponse)))))
.map(ResponseEntity::getBody)
.map(this::toChatResponse);
return response.switchMap(chatCompletionResponse -> {
if (this.isToolFunctionCall(chatCompletionResponse)) {
List<Message> toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(),
chatCompletionResponse);
return this.stream(new Prompt(toolCallMessageConversation, prompt.getOptions()));
}
return Mono.just(chatCompletionResponse).map(this::toChatResponse);
});
});
}
private List<Message> handleToolCallRequests(List<Message> previousMessages,
ChatCompletionResponse chatCompletionResponse) {
AnthropicMessage anthropicAssistantMessage = new AnthropicMessage(chatCompletionResponse.content(),
Role.ASSISTANT);
List<ContentBlock> toolToUseList = anthropicAssistantMessage.content()
.stream()
.filter(c -> c.type() == ContentBlock.ContentBlockType.TOOL_USE)
.toList();
List<AssistantMessage.ToolCall> toolCalls = new ArrayList<>();
for (ContentBlock toolToUse : toolToUseList) {
var functionCallId = toolToUse.id();
var functionName = toolToUse.name();
var functionArguments = ModelOptionsUtils.toJsonString(toolToUse.input());
toolCalls.add(new AssistantMessage.ToolCall(functionCallId, "function", functionName, functionArguments));
}
AssistantMessage assistantMessage = new AssistantMessage("", Map.of(), toolCalls);
ToolResponseMessage toolResponseMessage = this.executeFuncitons(assistantMessage);
// History
List<Message> toolCallMessageConversation = new ArrayList<>(previousMessages);
toolCallMessageConversation.add(assistantMessage);
toolCallMessageConversation.add(toolResponseMessage);
return toolCallMessageConversation;
}
private ChatResponse toChatResponse(ChatCompletionResponse chatCompletion) {
if (chatCompletion == null) {
logger.warn("Null chat completion returned");
@@ -203,18 +250,45 @@ public class AnthropicChatModel extends
List<AnthropicMessage> userMessages = prompt.getInstructions()
.stream()
.filter(m -> m.getMessageType() != MessageType.SYSTEM)
.map(m -> {
List<ContentBlock> contents = new ArrayList<>(List.of(new ContentBlock(m.getContent())));
if (!CollectionUtils.isEmpty(m.getMedia())) {
List<ContentBlock> mediaContent = m.getMedia()
.stream()
.map(media -> new ContentBlock(media.getMimeType().toString(),
this.fromMediaData(media.getData())))
.toList();
contents.addAll(mediaContent);
.filter(message -> message.getMessageType() != MessageType.SYSTEM)
.map(message -> {
if (message.getMessageType() == MessageType.USER) {
List<ContentBlock> contents = new ArrayList<>(List.of(new ContentBlock(message.getContent())));
if (!CollectionUtils.isEmpty(message.getMedia())) {
List<ContentBlock> mediaContent = message.getMedia()
.stream()
.map(media -> new ContentBlock(media.getMimeType().toString(),
this.fromMediaData(media.getData())))
.toList();
contents.addAll(mediaContent);
}
return new AnthropicMessage(contents, Role.valueOf(message.getMessageType().name()));
}
else if (message.getMessageType() == MessageType.ASSISTANT) {
AssistantMessage assistantMessage = (AssistantMessage) message;
List<ContentBlock> contentBlocks = new ArrayList<>();
if (StringUtils.hasText(message.getContent())) {
contentBlocks.add(new ContentBlock(message.getContent()));
}
if (!CollectionUtils.isEmpty(assistantMessage.getToolCalls())) {
for (AssistantMessage.ToolCall toolCall : assistantMessage.getToolCalls()) {
contentBlocks.add(new ContentBlock(ContentBlockType.TOOL_USE, toolCall.id(),
toolCall.name(), ModelOptionsUtils.jsonToMap(toolCall.arguments())));
}
}
return new AnthropicMessage(contentBlocks, Role.ASSISTANT);
}
else if (message.getMessageType() == MessageType.TOOL) {
List<ContentBlock> toolResponses = ((ToolResponseMessage) message).getResponses()
.stream()
.map(toolResponse -> new ContentBlock(ContentBlockType.TOOL_RESULT, toolResponse.id(),
toolResponse.responseData()))
.toList();
return new AnthropicMessage(toolResponses, Role.USER);
}
else {
throw new IllegalArgumentException("Unsupported message type: " + message.getMessageType());
}
return new AnthropicMessage(contents, Role.valueOf(m.getMessageType().name()));
})
.toList();
@@ -265,74 +339,17 @@ public class AnthropicChatModel extends
}).toList();
}
@Override
protected ChatCompletionRequest doCreateToolResponseRequest(ChatCompletionRequest previousRequest,
AnthropicMessage responseMessage, List<AnthropicMessage> conversationHistory) {
List<ContentBlock> toolToUseList = responseMessage.content()
.stream()
.filter(c -> c.type() == ContentBlock.ContentBlockType.TOOL_USE)
.toList();
List<ContentBlock> toolResults = new ArrayList<>();
for (ContentBlock toolToUse : toolToUseList) {
var functionCallId = toolToUse.id();
var functionName = toolToUse.name();
var functionArguments = toolToUse.input();
if (!this.functionCallbackRegister.containsKey(functionName)) {
throw new IllegalStateException("No function callback found for function name: " + functionName);
}
String functionResponse = this.functionCallbackRegister.get(functionName)
.call(ModelOptionsUtils.toJsonString(functionArguments));
toolResults.add(new ContentBlock(ContentBlockType.TOOL_RESULT, functionCallId, functionResponse));
}
// Add the function response to the conversation.
conversationHistory.add(new AnthropicMessage(toolResults, Role.USER));
// Recursively call chatCompletionWithTools until the model doesn't call a
// functions anymore.
return ChatCompletionRequest.from(previousRequest).withMessages(conversationHistory).build();
}
@Override
protected List<AnthropicMessage> doGetUserMessages(ChatCompletionRequest request) {
return request.messages();
}
@Override
protected AnthropicMessage doGetToolResponseMessage(ResponseEntity<ChatCompletionResponse> response) {
return new AnthropicMessage(response.getBody().content(), Role.ASSISTANT);
}
@Override
protected ResponseEntity<ChatCompletionResponse> doChatCompletion(ChatCompletionRequest request) {
return this.anthropicApi.chatCompletionEntity(request);
}
@SuppressWarnings("null")
@Override
protected boolean isToolFunctionCall(ResponseEntity<ChatCompletionResponse> response) {
if (response == null || response.getBody() == null || CollectionUtils.isEmpty(response.getBody().content())) {
protected boolean isToolFunctionCall(ChatCompletionResponse response) {
if (response == null || CollectionUtils.isEmpty(response.content())) {
return false;
}
return response.getBody()
.content()
return response.content()
.stream()
.anyMatch(content -> content.type() == ContentBlock.ContentBlockType.TOOL_USE);
}
@Override
protected Flux<ResponseEntity<ChatCompletionResponse>> doChatCompletionStream(ChatCompletionRequest request) {
return this.anthropicApi.chatCompletionStream(request).map(Optional::ofNullable).map(ResponseEntity::of);
}
@Override
public ChatOptions getDefaultOptions() {
return AnthropicChatOptions.fromOptions(this.defaultOptions);

View File

@@ -15,13 +15,20 @@
*/
package org.springframework.ai.openai;
import java.util.ArrayList;
import java.util.Base64;
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.MessageType;
import org.springframework.ai.chat.messages.ToolResponseMessage;
import org.springframework.ai.chat.messages.ToolResponseMessage.ToolResponse;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.model.ChatModel;
@@ -50,17 +57,10 @@ import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.MimeType;
import reactor.core.publisher.Flux;
import reactor.core.publisher.Mono;
import java.util.ArrayList;
import java.util.Base64;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
/**
* {@link ChatModel} and {@link StreamingChatModel} implementation for {@literal OpenAI}
* backed by {@link OpenAiApi}.
@@ -266,12 +266,12 @@ public class OpenAiChatModel extends AbstractToolCallSupport<ChatCompletion> imp
AssistantMessage assistantMessage = new AssistantMessage(nativeAssistantMessage.content(), Map.of(),
assistantToolCalls);
List<ToolResponseMessage> toolResponseMessages = this.executeFuncitons(assistantMessage, false);
ToolResponseMessage toolResponseMessage = this.executeFuncitons(assistantMessage);
// History
List<Message> messages = new ArrayList<>(previousMessages);
messages.add(assistantMessage);
messages.addAll(toolResponseMessages);
messages.add(toolResponseMessage);
return messages;
}
@@ -321,8 +321,8 @@ public class OpenAiChatModel extends AbstractToolCallSupport<ChatCompletion> imp
content = contentList;
}
return new ChatCompletionMessage(content,
ChatCompletionMessage.Role.valueOf(message.getMessageType().name()));
return List.of(new ChatCompletionMessage(content,
ChatCompletionMessage.Role.valueOf(message.getMessageType().name())));
}
else if (message.getMessageType() == MessageType.ASSISTANT) {
var assistantMessage = (AssistantMessage) message;
@@ -333,21 +333,27 @@ public class OpenAiChatModel extends AbstractToolCallSupport<ChatCompletion> imp
return new ToolCall(toolCall.id(), toolCall.type(), function);
}).toList();
}
return new ChatCompletionMessage(assistantMessage.getContent(), ChatCompletionMessage.Role.ASSISTANT,
null, null, toolCalls);
return List.of(new ChatCompletionMessage(assistantMessage.getContent(),
ChatCompletionMessage.Role.ASSISTANT, null, null, toolCalls));
}
else if (message.getMessageType() == MessageType.TOOL) {
ToolResponseMessage toolMessage = (ToolResponseMessage) message;
Assert.isTrue(toolMessage.getResponses().size() == 1,
"ToolResponseMessage must have exactly one response");
ToolResponse response = toolMessage.getResponses().get(0);
return new ChatCompletionMessage(response.respoinse(), ChatCompletionMessage.Role.TOOL, response.name(),
response.id(), null);
toolMessage.getResponses().forEach(response -> {
Assert.isTrue(response.id() != null, "ToolResponseMessage must have an id");
Assert.isTrue(response.name() != null, "ToolResponseMessage must have a name");
});
return toolMessage.getResponses()
.stream()
.map(tr -> new ChatCompletionMessage(tr.responseData(), ChatCompletionMessage.Role.TOOL, tr.name(),
tr.id(), null))
.toList();
}
else {
throw new IllegalArgumentException("Unsupported message type: " + message.getMessageType());
}
}).toList();
}).flatMap(List::stream).toList();
ChatCompletionRequest request = new ChatCompletionRequest(chatCompletionMessages, stream);

View File

@@ -198,12 +198,12 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport<GenerateCon
AssistantMessage assistantMessage = new AssistantMessage("", Map.of(), assistantToolCalls);
List<ToolResponseMessage> toolResponseMessages = this.executeFuncitons(assistantMessage, true);
ToolResponseMessage toolResponseMessage = this.executeFuncitons(assistantMessage);
// History
List<Message> toolCallMessageConversation = new ArrayList<>(previousMessages);
toolCallMessageConversation.add(assistantMessage);
toolCallMessageConversation.addAll(toolResponseMessages);
toolCallMessageConversation.add(toolResponseMessage);
return toolCallMessageConversation;
}
@@ -420,7 +420,7 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport<GenerateCon
.map(response -> Part.newBuilder()
.setFunctionResponse(FunctionResponse.newBuilder()
.setName(response.name())
.setResponse(jsonToStruct(response.respoinse()))
.setResponse(jsonToStruct(response.responseData()))
.build())
.build())
.toList();

View File

@@ -1,492 +0,0 @@
/*
* Copyright 2023 - 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vertexai.gemini;
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.Content;
import com.google.cloud.vertexai.api.Content.Builder;
import com.google.cloud.vertexai.api.FunctionCall;
import com.google.cloud.vertexai.api.FunctionDeclaration;
import com.google.cloud.vertexai.api.FunctionResponse;
import com.google.cloud.vertexai.api.GenerateContentResponse;
import com.google.cloud.vertexai.api.GenerationConfig;
import com.google.cloud.vertexai.api.Part;
import com.google.cloud.vertexai.api.Schema;
import com.google.cloud.vertexai.api.Tool;
import com.google.cloud.vertexai.generativeai.ContentMaker;
import com.google.cloud.vertexai.generativeai.GenerativeModel;
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.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
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.UserMessage;
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.model.function.AbstractFunctionCallSupport;
import org.springframework.ai.model.function.FunctionCallbackContext;
import org.springframework.ai.vertexai.gemini.metadata.VertexAiChatResponseMetadata;
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 java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Set;
import java.util.stream.Collectors;
/**
* @author Christian Tzolov
* @author Grogdunn
* @author luocongqiu
* @since 0.8.1
*/
public class VertexAiGeminiChatModelOld
extends AbstractFunctionCallSupport<Content, VertexAiGeminiChatModelOld.GeminiRequest, GenerateContentResponse>
implements ChatModel, DisposableBean {
private final static boolean IS_RUNTIME_CALL = true;
private final VertexAI vertexAI;
private final VertexAiGeminiChatOptions defaultOptions;
private final GenerationConfig generationConfig;
public enum GeminiMessageType {
USER("user"),
MODEL("model");
GeminiMessageType(String value) {
this.value = value;
}
public final String value;
public String getValue() {
return this.value;
}
}
public enum ChatModel implements ChatModelDescription {
GEMINI_PRO_VISION("gemini-pro-vision"),
GEMINI_PRO("gemini-pro"),
GEMINI_1_5_PRO("gemini-1.5-pro-001"),
GEMINI_1_5_FLASH("gemini-1.5-flash-001");
ChatModel(String value) {
this.value = value;
}
public final String value;
public String getValue() {
return this.value;
}
@Override
public String getName() {
return this.value;
}
}
public VertexAiGeminiChatModelOld(VertexAI vertexAI) {
this(vertexAI, VertexAiGeminiChatOptions.builder()
// .withModel(VertexAiGeminiChatModelOld.ChatModel.GEMINI_PRO_VISION)
.withTemperature(0.8f)
.build());
}
public VertexAiGeminiChatModelOld(VertexAI vertexAI, VertexAiGeminiChatOptions options) {
this(vertexAI, options, null);
}
public VertexAiGeminiChatModelOld(VertexAI vertexAI, VertexAiGeminiChatOptions options,
FunctionCallbackContext functionCallbackContext) {
super(functionCallbackContext);
Assert.notNull(vertexAI, "VertexAI must not be null");
Assert.notNull(options, "VertexAiGeminiChatOptions must not be null");
Assert.notNull(options.getModel(), "VertexAiGeminiChatOptions.modelName must not be null");
this.vertexAI = vertexAI;
this.defaultOptions = options;
this.generationConfig = toGenerationConfig(options);
}
// https://cloud.google.com/vertex-ai/docs/generative-ai/model-reference/gemini
@Override
public ChatResponse call(Prompt prompt) {
var geminiRequest = createGeminiRequest(prompt);
GenerateContentResponse response = this.callWithFunctionSupport(geminiRequest);
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));
}
@Override
public Flux<ChatResponse> stream(Prompt prompt) {
try {
var request = createGeminiRequest(prompt);
ResponseStream<GenerateContentResponse> responseStream = request.model
.generateContentStream(request.contents);
return Flux.fromStream(responseStream.stream())
.switchMap(r -> handleFunctionCallOrReturnStream(request, Flux.just(r)))
.map(response -> {
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));
});
}
catch (Exception e) {
throw new RuntimeException("Failed to generate content", e);
}
}
private VertexAiChatResponseMetadata toChatResponseMetadata(GenerateContentResponse response) {
return new VertexAiChatResponseMetadata(new VertexAiUsage(response.getUsageMetadata()));
}
@JsonInclude(Include.NON_NULL)
public record GeminiRequest(List<Content> contents, GenerativeModel model) {
}
private GeminiRequest createGeminiRequest(Prompt prompt) {
Set<String> functionsForThisRequest = new HashSet<>();
GenerationConfig generationConfig = this.generationConfig;
var generativeModelBuilder = new GenerativeModel.Builder().setModelName(this.defaultOptions.getModel())
.setVertexAi(this.vertexAI);
VertexAiGeminiChatOptions updatedRuntimeOptions = null;
if (prompt.getOptions() != null) {
updatedRuntimeOptions = ModelOptionsUtils.copyToTarget(prompt.getOptions(), ChatOptions.class,
VertexAiGeminiChatOptions.class);
functionsForThisRequest
.addAll(handleFunctionCallbackConfigurations(updatedRuntimeOptions, IS_RUNTIME_CALL));
}
if (this.defaultOptions != null) {
functionsForThisRequest.addAll(handleFunctionCallbackConfigurations(this.defaultOptions, !IS_RUNTIME_CALL));
if (updatedRuntimeOptions == null) {
updatedRuntimeOptions = VertexAiGeminiChatOptions.builder().build();
}
updatedRuntimeOptions = ModelOptionsUtils.merge(updatedRuntimeOptions, this.defaultOptions,
VertexAiGeminiChatOptions.class);
}
if (updatedRuntimeOptions != null) {
if (StringUtils.hasText(updatedRuntimeOptions.getModel())
&& !updatedRuntimeOptions.getModel().equals(this.defaultOptions.getModel())) {
// Override model name
generativeModelBuilder.setModelName(updatedRuntimeOptions.getModel());
}
generationConfig = toGenerationConfig(updatedRuntimeOptions);
}
// Add the enabled functions definitions to the request's tools parameter.
if (!CollectionUtils.isEmpty(functionsForThisRequest)) {
List<Tool> tools = this.getFunctionTools(functionsForThisRequest);
generativeModelBuilder.setTools(tools);
}
generativeModelBuilder.setGenerationConfig(generationConfig);
GenerativeModel generativeModel = generativeModelBuilder.build();
String systemContext = prompt.getInstructions()
.stream()
.filter(m -> m.getMessageType() == MessageType.SYSTEM)
.map(m -> m.getContent())
.collect(Collectors.joining(System.lineSeparator()));
if (StringUtils.hasText(systemContext)) {
generativeModel.withSystemInstruction(ContentMaker.fromString(systemContext));
}
return new GeminiRequest(toGeminiContent(prompt), generativeModel);
}
private GenerationConfig toGenerationConfig(VertexAiGeminiChatOptions options) {
GenerationConfig.Builder generationConfigBuilder = GenerationConfig.newBuilder();
if (options.getTemperature() != null) {
generationConfigBuilder.setTemperature(options.getTemperature());
}
if (options.getMaxOutputTokens() != null) {
generationConfigBuilder.setMaxOutputTokens(options.getMaxOutputTokens());
}
if (options.getTopK() != null) {
generationConfigBuilder.setTopK(options.getTopK());
}
if (options.getTopP() != null) {
generationConfigBuilder.setTopP(options.getTopP());
}
if (options.getCandidateCount() != null) {
generationConfigBuilder.setCandidateCount(options.getCandidateCount());
}
if (options.getStopSequences() != null) {
generationConfigBuilder.addAllStopSequences(options.getStopSequences());
}
return generationConfigBuilder.build();
}
private List<Content> toGeminiContent(Prompt prompt) {
List<Content> contents = prompt.getInstructions()
.stream()
.filter(m -> m.getMessageType() == MessageType.USER || m.getMessageType() == MessageType.ASSISTANT)
.map(message -> Content.newBuilder()
.setRole(toGeminiMessageType(message.getMessageType()).getValue())
.addAllParts(messageToGeminiParts(message))
.build())
.toList();
return contents;
}
private static GeminiMessageType toGeminiMessageType(@NonNull MessageType type) {
Assert.notNull(type, "Message type must not be null");
switch (type) {
case USER:
return GeminiMessageType.USER;
case ASSISTANT:
return GeminiMessageType.MODEL;
default:
throw new IllegalArgumentException("Unsupported message type: " + type);
}
}
static List<Part> messageToGeminiParts(Message message) {
if (message instanceof UserMessage userMessage) {
String messageTextContent = (userMessage.getContent() == null) ? "null" : userMessage.getContent();
Part textPart = Part.newBuilder().setText(messageTextContent).build();
List<Part> parts = new ArrayList<>(List.of(textPart));
List<Part> mediaParts = userMessage.getMedia()
.stream()
.map(mediaData -> PartMaker.fromMimeTypeAndData(mediaData.getMimeType().toString(),
mediaData.getData()))
.toList();
if (!CollectionUtils.isEmpty(mediaParts)) {
parts.addAll(mediaParts);
}
return parts;
}
else if (message instanceof AssistantMessage assistantMessage) {
return List.of(Part.newBuilder().setText(assistantMessage.getContent()).build());
}
else {
throw new IllegalArgumentException("Gemini doesn't support message type: " + message.getClass());
}
}
private List<Tool> getFunctionTools(Set<String> functionNames) {
final var tool = Tool.newBuilder();
final List<FunctionDeclaration> functionDeclarations = this.resolveFunctionCallbacks(functionNames)
.stream()
.map(functionCallback -> FunctionDeclaration.newBuilder()
.setName(functionCallback.getName())
.setDescription(functionCallback.getDescription())
.setParameters(jsonToSchema(functionCallback.getInputTypeSchema()))
.build())
.toList();
tool.addAllFunctionDeclarations(functionDeclarations);
return List.of(tool.build());
}
private static String structToJson(Struct struct) {
try {
return JsonFormat.printer().print(struct);
}
catch (Exception e) {
throw new RuntimeException(e);
}
}
private static Struct jsonToStruct(String json) {
try {
var structBuilder = Struct.newBuilder();
JsonFormat.parser().ignoringUnknownFields().merge(json, structBuilder);
return structBuilder.build();
}
catch (Exception e) {
throw new RuntimeException(e);
}
}
private static Schema jsonToSchema(String json) {
try {
var schemaBuilder = Schema.newBuilder();
JsonFormat.parser().ignoringUnknownFields().merge(json, schemaBuilder);
return schemaBuilder.build();
}
catch (Exception e) {
throw new RuntimeException(e);
}
}
@Override
public void destroy() throws Exception {
if (this.vertexAI != null) {
this.vertexAI.close();
}
}
@Override
protected GeminiRequest doCreateToolResponseRequest(GeminiRequest previousRequest, Content responseMessage,
List<Content> conversationHistory) {
var iterator = responseMessage.getPartsList().iterator();
Builder builder = Content.newBuilder();
while (iterator.hasNext()) {
FunctionCall functionCall = iterator.next().getFunctionCall();
var functionName = functionCall.getName();
String functionArguments = structToJson(functionCall.getArgs());
if (!this.functionCallbackRegister.containsKey(functionName)) {
throw new IllegalStateException("No function callback found for function name: " + functionName);
}
String functionResponse = this.functionCallbackRegister.get(functionName).call(functionArguments);
builder.addParts(Part.newBuilder()
.setFunctionResponse(FunctionResponse.newBuilder()
.setName(functionCall.getName())
.setResponse(jsonToStruct(functionResponse))
.build())
.build());
}
conversationHistory.add(builder.build());
return new GeminiRequest(conversationHistory, previousRequest.model());
}
@Override
protected List<Content> doGetUserMessages(GeminiRequest request) {
return request.contents;
}
@Override
protected Content doGetToolResponseMessage(GenerateContentResponse response) {
return response.getCandidatesList().get(0).getContent();
}
@Override
protected GenerateContentResponse doChatCompletion(GeminiRequest request) {
try {
return request.model.generateContent(request.contents);
}
catch (Exception e) {
throw new RuntimeException("Failed to generate content", e);
}
}
@Override
protected Flux<GenerateContentResponse> doChatCompletionStream(GeminiRequest request) {
try {
ResponseStream<GenerateContentResponse> responseStream = request.model
.generateContentStream(request.contents);
return Flux.fromStream(responseStream.stream());
}
catch (Exception e) {
throw new RuntimeException("Failed to generate content", e);
}
}
@Override
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);
}
}

View File

@@ -29,7 +29,7 @@ import java.util.Objects;
*/
public class ToolResponseMessage extends AbstractMessage {
public record ToolResponse(String id, String name, String respoinse) {
public record ToolResponse(String id, String name, String responseData) {
};
private List<ToolResponse> responses = new ArrayList<>();

View File

@@ -127,9 +127,7 @@ public abstract class AbstractToolCallSupport<TRes> {
return retrievedFunctionCallbacks;
}
protected List<ToolResponseMessage> executeFuncitons(AssistantMessage assistantMessage, boolean signelResponse) {
List<ToolResponseMessage> toolResponseMessages = new ArrayList<>();
protected ToolResponseMessage executeFuncitons(AssistantMessage assistantMessage) {
List<ToolResponseMessage.ToolResponse> toolResponses = new ArrayList<>();
@@ -147,15 +145,7 @@ public abstract class AbstractToolCallSupport<TRes> {
toolResponses.add(new ToolResponseMessage.ToolResponse(toolCall.id(), functionName, functionResponse));
}
if (signelResponse) {
toolResponseMessages.add(new ToolResponseMessage(toolResponses, Map.of()));
}
else {
for (ToolResponseMessage.ToolResponse toolResponse : toolResponses) {
toolResponseMessages.add(new ToolResponseMessage(List.of(toolResponse)));
}
}
return toolResponseMessages;
return new ToolResponseMessage(toolResponses, Map.of());
}
abstract protected boolean isToolFunctionCall(TRes response);