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:
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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<>();
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user