High-level API function calling support for VertexAI Gemini
- Refactor VertexAiChatModel's function calling handling to use Spring AI abstractions.
This commit is contained in:
@@ -21,6 +21,7 @@ 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;
|
||||
@@ -265,7 +266,7 @@ public class OpenAiChatModel extends AbstractToolCallSupport<ChatCompletion> imp
|
||||
AssistantMessage assistantMessage = new AssistantMessage(nativeAssistantMessage.content(), Map.of(),
|
||||
assistantToolCalls);
|
||||
|
||||
List<ToolResponseMessage> toolResponseMessages = this.executeFuncitons(assistantMessage);
|
||||
List<ToolResponseMessage> toolResponseMessages = this.executeFuncitons(assistantMessage, false);
|
||||
|
||||
// History
|
||||
List<Message> messages = new ArrayList<>(previousMessages);
|
||||
@@ -337,8 +338,11 @@ public class OpenAiChatModel extends AbstractToolCallSupport<ChatCompletion> imp
|
||||
}
|
||||
else if (message.getMessageType() == MessageType.TOOL) {
|
||||
ToolResponseMessage toolMessage = (ToolResponseMessage) message;
|
||||
return new ChatCompletionMessage(toolMessage.getContent(), ChatCompletionMessage.Role.TOOL,
|
||||
toolMessage.getName(), toolMessage.getId(), null);
|
||||
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);
|
||||
}
|
||||
else {
|
||||
throw new IllegalArgumentException("Unsupported message type: " + message.getMessageType());
|
||||
|
||||
@@ -43,11 +43,11 @@ import java.util.stream.Collectors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
@SpringBootTest(classes = OpenAiChatModel3IT.Config.class)
|
||||
@SpringBootTest(classes = OpenAiChatModelFunctionCallingIT.Config.class)
|
||||
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
|
||||
class OpenAiChatModel3IT {
|
||||
class OpenAiChatModelFunctionCallingIT {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(OpenAiChatModel3IT.class);
|
||||
private static final Logger logger = LoggerFactory.getLogger(OpenAiChatModelFunctionCallingIT.class);
|
||||
|
||||
@Autowired
|
||||
ChatModel chatModel;
|
||||
@@ -72,9 +72,7 @@ class OpenAiChatModel3IT {
|
||||
|
||||
logger.info("Response: {}", response);
|
||||
|
||||
assertThat(response.getResult().getOutput().getContent()).containsAnyOf("30.0", "30");
|
||||
assertThat(response.getResult().getOutput().getContent()).containsAnyOf("10.0", "10");
|
||||
assertThat(response.getResult().getOutput().getContent()).containsAnyOf("15.0", "15");
|
||||
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -105,9 +103,7 @@ class OpenAiChatModel3IT {
|
||||
.collect(Collectors.joining());
|
||||
logger.info("Response: {}", content);
|
||||
|
||||
assertThat(content).containsAnyOf("30.0", "30");
|
||||
assertThat(content).containsAnyOf("10.0", "10");
|
||||
assertThat(content).containsAnyOf("15.0", "15");
|
||||
assertThat(content).contains("30", "10", "15");
|
||||
}
|
||||
|
||||
@SpringBootConfiguration
|
||||
@@ -15,11 +15,39 @@
|
||||
*/
|
||||
package org.springframework.ai.vertexai.gemini;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
import java.util.stream.Collectors;
|
||||
|
||||
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.UserMessage;
|
||||
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.model.function.AbstractToolCallSupport;
|
||||
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 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;
|
||||
@@ -34,33 +62,9 @@ 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;
|
||||
import reactor.core.publisher.Flux;
|
||||
import reactor.core.publisher.Mono;
|
||||
|
||||
/**
|
||||
* @author Christian Tzolov
|
||||
@@ -68,8 +72,7 @@ import java.util.stream.Collectors;
|
||||
* @author luocongqiu
|
||||
* @since 0.8.1
|
||||
*/
|
||||
public class VertexAiGeminiChatModel
|
||||
extends AbstractFunctionCallSupport<Content, VertexAiGeminiChatModel.GeminiRequest, GenerateContentResponse>
|
||||
public class VertexAiGeminiChatModel extends AbstractToolCallSupport<GenerateContentResponse>
|
||||
implements ChatModel, DisposableBean {
|
||||
|
||||
private final static boolean IS_RUNTIME_CALL = true;
|
||||
@@ -157,7 +160,15 @@ public class VertexAiGeminiChatModel
|
||||
|
||||
var geminiRequest = createGeminiRequest(prompt);
|
||||
|
||||
GenerateContentResponse response = this.callWithFunctionSupport(geminiRequest);
|
||||
GenerateContentResponse response = this.getContentResponse(geminiRequest);
|
||||
|
||||
// GenerateContentResponse response = this.callWithFunctionSupport(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()
|
||||
@@ -170,6 +181,32 @@ public class VertexAiGeminiChatModel
|
||||
return new ChatResponse(generations, toChatResponseMetadata(response));
|
||||
}
|
||||
|
||||
public List<Message> handleToolCallRequests(List<Message> previousMessages, GenerateContentResponse response) {
|
||||
|
||||
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);
|
||||
|
||||
List<ToolResponseMessage> toolResponseMessages = this.executeFuncitons(assistantMessage, true);
|
||||
|
||||
// History
|
||||
List<Message> toolCallMessageConversation = new ArrayList<>(previousMessages);
|
||||
toolCallMessageConversation.add(assistantMessage);
|
||||
toolCallMessageConversation.addAll(toolResponseMessages);
|
||||
return toolCallMessageConversation;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatResponse> stream(Prompt prompt) {
|
||||
try {
|
||||
@@ -179,9 +216,16 @@ public class VertexAiGeminiChatModel
|
||||
ResponseStream<GenerateContentResponse> responseStream = request.model
|
||||
.generateContentStream(request.contents);
|
||||
|
||||
return Flux.fromStream(responseStream.stream())
|
||||
.switchMap(r -> handleFunctionCallOrReturnStream(request, Flux.just(r)))
|
||||
.map(response -> {
|
||||
return Flux.fromStream(responseStream.stream()).switchMap(response -> {
|
||||
if (this.isToolFunctionCall(response)) {
|
||||
List<Message> toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(),
|
||||
response);
|
||||
// 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 Mono.just(response).map(response2 -> {
|
||||
List<Generation> generations = response.getCandidatesList()
|
||||
.stream()
|
||||
.map(candidate -> candidate.getContent().getPartsList())
|
||||
@@ -191,7 +235,9 @@ public class VertexAiGeminiChatModel
|
||||
.toList();
|
||||
|
||||
return new ChatResponse(generations, toChatResponseMetadata(response));
|
||||
|
||||
});
|
||||
});
|
||||
}
|
||||
catch (Exception e) {
|
||||
throw new RuntimeException("Failed to generate content", e);
|
||||
@@ -302,7 +348,8 @@ public class VertexAiGeminiChatModel
|
||||
|
||||
List<Content> contents = prompt.getInstructions()
|
||||
.stream()
|
||||
.filter(m -> m.getMessageType() == MessageType.USER || m.getMessageType() == MessageType.ASSISTANT)
|
||||
.filter(m -> m.getMessageType() == MessageType.USER || m.getMessageType() == MessageType.ASSISTANT
|
||||
|| m.getMessageType() == MessageType.TOOL)
|
||||
.map(message -> Content.newBuilder()
|
||||
.setRole(toGeminiMessageType(message.getMessageType()).getValue())
|
||||
.addAllParts(messageToGeminiParts(message))
|
||||
@@ -318,6 +365,7 @@ public class VertexAiGeminiChatModel
|
||||
|
||||
switch (type) {
|
||||
case USER:
|
||||
case TOOL:
|
||||
return GeminiMessageType.USER;
|
||||
case ASSISTANT:
|
||||
return GeminiMessageType.MODEL;
|
||||
@@ -348,7 +396,34 @@ public class VertexAiGeminiChatModel
|
||||
return parts;
|
||||
}
|
||||
else if (message instanceof AssistantMessage assistantMessage) {
|
||||
return List.of(Part.newBuilder().setText(assistantMessage.getContent()).build());
|
||||
List<Part> parts = new ArrayList<>();
|
||||
if (StringUtils.hasText(assistantMessage.getContent())) {
|
||||
List.of(Part.newBuilder().setText(assistantMessage.getContent()).build());
|
||||
}
|
||||
if (!CollectionUtils.isEmpty(assistantMessage.getToolCalls())) {
|
||||
parts.addAll(assistantMessage.getToolCalls()
|
||||
.stream()
|
||||
.map(toolCall -> Part.newBuilder()
|
||||
.setFunctionCall(FunctionCall.newBuilder()
|
||||
.setName(toolCall.name())
|
||||
.setArgs(jsonToStruct(toolCall.arguments()))
|
||||
.build())
|
||||
.build())
|
||||
.toList());
|
||||
}
|
||||
return parts;
|
||||
}
|
||||
else if (message instanceof ToolResponseMessage toolResponseMessage) {
|
||||
|
||||
return toolResponseMessage.getResponses()
|
||||
.stream()
|
||||
.map(response -> Part.newBuilder()
|
||||
.setFunctionResponse(FunctionResponse.newBuilder()
|
||||
.setName(response.name())
|
||||
.setResponse(jsonToStruct(response.respoinse()))
|
||||
.build())
|
||||
.build())
|
||||
.toList();
|
||||
}
|
||||
else {
|
||||
throw new IllegalArgumentException("Gemini doesn't support message type: " + message.getClass());
|
||||
@@ -402,58 +477,7 @@ public class VertexAiGeminiChatModel
|
||||
}
|
||||
}
|
||||
|
||||
@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) {
|
||||
private GenerateContentResponse getContentResponse(GeminiRequest request) {
|
||||
try {
|
||||
return request.model.generateContent(request.contents);
|
||||
}
|
||||
@@ -462,19 +486,6 @@ public class VertexAiGeminiChatModel
|
||||
}
|
||||
}
|
||||
|
||||
@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())
|
||||
@@ -490,4 +501,11 @@ public class VertexAiGeminiChatModel
|
||||
return VertexAiGeminiChatOptions.fromOptions(this.defaultOptions);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void destroy() throws Exception {
|
||||
if (this.vertexAI != null) {
|
||||
this.vertexAI.close();
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -0,0 +1,492 @@
|
||||
/*
|
||||
* 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);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -136,6 +136,45 @@ public class VertexAiGeminiChatModelFunctionCallingIT {
|
||||
|
||||
}
|
||||
|
||||
@Test
|
||||
public void functionCallTestInferredOpenApiSchema2() {
|
||||
|
||||
UserMessage userMessage = new UserMessage(
|
||||
"What's the weather like in San Francisco, Paris and in Tokyo? Return the temperature in Celsius.");
|
||||
|
||||
List<Message> messages = new ArrayList<>(List.of(userMessage));
|
||||
|
||||
var promptOptions = VertexAiGeminiChatOptions.builder()
|
||||
.withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_1_5_FLASH)
|
||||
.withFunctionCallbacks(List.of(
|
||||
FunctionCallbackWrapper.builder(new MockWeatherService())
|
||||
.withSchemaType(SchemaType.OPEN_API_SCHEMA)
|
||||
.withName("get_current_weather")
|
||||
.withDescription("Get the current weather in a given location.")
|
||||
.build(),
|
||||
FunctionCallbackWrapper.builder(new PaymentStatus())
|
||||
.withSchemaType(SchemaType.OPEN_API_SCHEMA)
|
||||
.withName("get_payment_status")
|
||||
.withDescription(
|
||||
"Retrieves the payment status for transaction. For example what is the payment status for transaction 700?")
|
||||
.build()))
|
||||
.build();
|
||||
|
||||
ChatResponse response = chatModel.call(new Prompt(messages, promptOptions));
|
||||
|
||||
logger.info("Response: {}", response);
|
||||
|
||||
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
|
||||
|
||||
ChatResponse response2 = chatModel
|
||||
.call(new Prompt("What is the payment status for transaction 696?", promptOptions));
|
||||
|
||||
logger.info("Response: {}", response2);
|
||||
|
||||
assertThat(response2.getResult().getOutput().getContent()).containsIgnoringCase("transaction 696 is PAYED");
|
||||
|
||||
}
|
||||
|
||||
@Test
|
||||
public void functionCallTestInferredOpenApiSchemaStream() {
|
||||
|
||||
|
||||
@@ -15,6 +15,8 @@
|
||||
*/
|
||||
package org.springframework.ai.chat.messages;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Objects;
|
||||
|
||||
@@ -27,31 +29,27 @@ import java.util.Objects;
|
||||
*/
|
||||
public class ToolResponseMessage extends AbstractMessage {
|
||||
|
||||
private final String id;
|
||||
public record ToolResponse(String id, String name, String respoinse) {
|
||||
};
|
||||
|
||||
private final String name;
|
||||
private List<ToolResponse> responses = new ArrayList<>();
|
||||
|
||||
public ToolResponseMessage(String id, String name, String content) {
|
||||
this(id, name, content, Map.of());
|
||||
public ToolResponseMessage(List<ToolResponse> responses) {
|
||||
this(responses, Map.of());
|
||||
}
|
||||
|
||||
public ToolResponseMessage(String id, String name, String content, Map<String, Object> metadata) {
|
||||
super(MessageType.TOOL, content, metadata);
|
||||
this.id = id;
|
||||
this.name = name;
|
||||
public ToolResponseMessage(List<ToolResponse> responses, Map<String, Object> metadata) {
|
||||
super(MessageType.TOOL, "", metadata);
|
||||
this.responses = responses;
|
||||
}
|
||||
|
||||
public String getId() {
|
||||
return id;
|
||||
}
|
||||
|
||||
public String getName() {
|
||||
return name;
|
||||
public List<ToolResponse> getResponses() {
|
||||
return this.responses;
|
||||
}
|
||||
|
||||
@Override
|
||||
public int hashCode() {
|
||||
return Objects.hash(this.id, this.name, getContent(), this.metadata, this.messageType);
|
||||
return Objects.hash(this.responses, getContent(), this.metadata, this.messageType);
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -63,15 +61,14 @@ public class ToolResponseMessage extends AbstractMessage {
|
||||
return false;
|
||||
}
|
||||
ToolResponseMessage other = (ToolResponseMessage) obj;
|
||||
return Objects.equals(id, other.id) && Objects.equals(this.name, other.name)
|
||||
&& Objects.equals(getContent(), other.getContent()) && Objects.equals(this.metadata, other.metadata)
|
||||
&& this.messageType == other.messageType;
|
||||
return Objects.equals(this.responses, other.responses) && Objects.equals(getContent(), other.getContent())
|
||||
&& Objects.equals(this.metadata, other.metadata) && this.messageType == other.messageType;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String toString() {
|
||||
return "FunctionMessage [id=" + id + ", name=" + name + ", messageType=" + messageType + ", textContent="
|
||||
+ textContent + "]";
|
||||
return "ToolResponseMessage [responses=" + responses + ", messageType=" + messageType + ", metadata=" + metadata
|
||||
+ "]";
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -17,6 +17,7 @@ package org.springframework.ai.chat.prompt;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Objects;
|
||||
|
||||
@@ -120,8 +121,8 @@ public class Prompt implements ModelRequest<List<Message>> {
|
||||
assistantMessage.getToolCalls()));
|
||||
}
|
||||
else if (message instanceof ToolResponseMessage toolResponseMessage) {
|
||||
messagesCopy.add(new ToolResponseMessage(toolResponseMessage.getId(), toolResponseMessage.getName(),
|
||||
toolResponseMessage.getContent(), toolResponseMessage.getMetadata()));
|
||||
messagesCopy.add(new ToolResponseMessage(new ArrayList<>(toolResponseMessage.getResponses()),
|
||||
new HashMap<>(toolResponseMessage.getMetadata())));
|
||||
}
|
||||
else {
|
||||
throw new IllegalArgumentException("Unsupported message type: " + message.getClass().getName());
|
||||
|
||||
@@ -127,10 +127,12 @@ public abstract class AbstractToolCallSupport<TRes> {
|
||||
return retrievedFunctionCallbacks;
|
||||
}
|
||||
|
||||
protected List<ToolResponseMessage> executeFuncitons(AssistantMessage assistantMessage) {
|
||||
protected List<ToolResponseMessage> executeFuncitons(AssistantMessage assistantMessage, boolean signelResponse) {
|
||||
|
||||
List<ToolResponseMessage> toolResponseMessages = new ArrayList<>();
|
||||
|
||||
List<ToolResponseMessage.ToolResponse> toolResponses = new ArrayList<>();
|
||||
|
||||
for (AssistantMessage.ToolCall toolCall : assistantMessage.getToolCalls()) {
|
||||
|
||||
var functionName = toolCall.name();
|
||||
@@ -142,11 +144,18 @@ public abstract class AbstractToolCallSupport<TRes> {
|
||||
|
||||
String functionResponse = this.functionCallbackRegister.get(functionName).call(functionArguments);
|
||||
|
||||
toolResponseMessages.add(new ToolResponseMessage(toolCall.id(), functionName, functionResponse));
|
||||
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;
|
||||
|
||||
}
|
||||
|
||||
abstract protected boolean isToolFunctionCall(TRes response);
|
||||
|
||||
@@ -15,17 +15,17 @@
|
||||
*/
|
||||
package org.springframework.ai.autoconfigure.vertexai.gemini.tool;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import org.springframework.ai.autoconfigure.vertexai.gemini.VertexAiGeminiAutoConfiguration;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.messages.SystemMessage;
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.function.FunctionCallbackWrapper;
|
||||
import org.springframework.ai.model.function.FunctionCallbackWrapper.Builder.SchemaType;
|
||||
@@ -34,8 +34,6 @@ import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatOptions;
|
||||
import org.springframework.boot.autoconfigure.AutoConfigurations;
|
||||
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
@EnabledIfEnvironmentVariable(named = "VERTEX_AI_GEMINI_PROJECT_ID", matches = ".*")
|
||||
@EnabledIfEnvironmentVariable(named = "VERTEX_AI_GEMINI_LOCATION", matches = ".*")
|
||||
public class FunctionCallWithPromptFunctionIT {
|
||||
|
||||
Reference in New Issue
Block a user