diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java index cbe8ea2a6..35227909b 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java @@ -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> - implements ChatModel { +public class AnthropicChatModel extends AbstractToolCallSupport 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 completionEntity = this.callWithFunctionSupport(request); + ResponseEntity completionEntity = this.anthropicApi.chatCompletionEntity(request); + + if (this.isToolFunctionCall(completionEntity.getBody())) { + List 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 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 toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(), + chatCompletionResponse); + return this.stream(new Prompt(toolCallMessageConversation, prompt.getOptions())); + } + + return Mono.just(chatCompletionResponse).map(this::toChatResponse); + }); }); } + private List handleToolCallRequests(List previousMessages, + ChatCompletionResponse chatCompletionResponse) { + + AnthropicMessage anthropicAssistantMessage = new AnthropicMessage(chatCompletionResponse.content(), + Role.ASSISTANT); + + List toolToUseList = anthropicAssistantMessage.content() + .stream() + .filter(c -> c.type() == ContentBlock.ContentBlockType.TOOL_USE) + .toList(); + + List 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 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 userMessages = prompt.getInstructions() .stream() - .filter(m -> m.getMessageType() != MessageType.SYSTEM) - .map(m -> { - List contents = new ArrayList<>(List.of(new ContentBlock(m.getContent()))); - if (!CollectionUtils.isEmpty(m.getMedia())) { - List 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 contents = new ArrayList<>(List.of(new ContentBlock(message.getContent()))); + if (!CollectionUtils.isEmpty(message.getMedia())) { + List 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 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 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 conversationHistory) { - - List toolToUseList = responseMessage.content() - .stream() - .filter(c -> c.type() == ContentBlock.ContentBlockType.TOOL_USE) - .toList(); - - List 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 doGetUserMessages(ChatCompletionRequest request) { - return request.messages(); - } - - @Override - protected AnthropicMessage doGetToolResponseMessage(ResponseEntity response) { - return new AnthropicMessage(response.getBody().content(), Role.ASSISTANT); - } - - @Override - protected ResponseEntity doChatCompletion(ChatCompletionRequest request) { - return this.anthropicApi.chatCompletionEntity(request); - } - @SuppressWarnings("null") @Override - protected boolean isToolFunctionCall(ResponseEntity 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> doChatCompletionStream(ChatCompletionRequest request) { - - return this.anthropicApi.chatCompletionStream(request).map(Optional::ofNullable).map(ResponseEntity::of); - } - @Override public ChatOptions getDefaultOptions() { return AnthropicChatOptions.fromOptions(this.defaultOptions); diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java index 54ccd5c6b..fca2de714 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java @@ -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 imp AssistantMessage assistantMessage = new AssistantMessage(nativeAssistantMessage.content(), Map.of(), assistantToolCalls); - List toolResponseMessages = this.executeFuncitons(assistantMessage, false); + ToolResponseMessage toolResponseMessage = this.executeFuncitons(assistantMessage); // History List messages = new ArrayList<>(previousMessages); messages.add(assistantMessage); - messages.addAll(toolResponseMessages); + messages.add(toolResponseMessage); return messages; } @@ -321,8 +321,8 @@ public class OpenAiChatModel extends AbstractToolCallSupport 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 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); diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java index 457c48bfe..3ecd72c2e 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java @@ -198,12 +198,12 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport toolResponseMessages = this.executeFuncitons(assistantMessage, true); + ToolResponseMessage toolResponseMessage = this.executeFuncitons(assistantMessage); // History List toolCallMessageConversation = new ArrayList<>(previousMessages); toolCallMessageConversation.add(assistantMessage); - toolCallMessageConversation.addAll(toolResponseMessages); + toolCallMessageConversation.add(toolResponseMessage); return toolCallMessageConversation; } @@ -420,7 +420,7 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport Part.newBuilder() .setFunctionResponse(FunctionResponse.newBuilder() .setName(response.name()) - .setResponse(jsonToStruct(response.respoinse())) + .setResponse(jsonToStruct(response.responseData())) .build()) .build()) .toList(); diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModelOld.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModelOld.java deleted file mode 100644 index 0b65afcce..000000000 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModelOld.java +++ /dev/null @@ -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 - 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 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 stream(Prompt prompt) { - try { - - var request = createGeminiRequest(prompt); - - ResponseStream responseStream = request.model - .generateContentStream(request.contents); - - return Flux.fromStream(responseStream.stream()) - .switchMap(r -> handleFunctionCallOrReturnStream(request, Flux.just(r))) - .map(response -> { - List 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 contents, GenerativeModel model) { - } - - private GeminiRequest createGeminiRequest(Prompt prompt) { - - Set 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 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 toGeminiContent(Prompt prompt) { - - List 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 messageToGeminiParts(Message message) { - - if (message instanceof UserMessage userMessage) { - - String messageTextContent = (userMessage.getContent() == null) ? "null" : userMessage.getContent(); - Part textPart = Part.newBuilder().setText(messageTextContent).build(); - - List parts = new ArrayList<>(List.of(textPart)); - - List 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 getFunctionTools(Set functionNames) { - - final var tool = Tool.newBuilder(); - - final List 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 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 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 doChatCompletionStream(GeminiRequest request) { - try { - ResponseStream 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); - } - -} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/ToolResponseMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/ToolResponseMessage.java index 85fec0b22..b5114b7f7 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/ToolResponseMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/ToolResponseMessage.java @@ -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 responses = new ArrayList<>(); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java index 146e42df6..db684c65d 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java @@ -127,9 +127,7 @@ public abstract class AbstractToolCallSupport { return retrievedFunctionCallbacks; } - protected List executeFuncitons(AssistantMessage assistantMessage, boolean signelResponse) { - - List toolResponseMessages = new ArrayList<>(); + protected ToolResponseMessage executeFuncitons(AssistantMessage assistantMessage) { List toolResponses = new ArrayList<>(); @@ -147,15 +145,7 @@ public abstract class AbstractToolCallSupport { 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);