diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/MessageToPromptConverter.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/MessageToPromptConverter.java index 6f2f3301e..95abde872 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/MessageToPromptConverter.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/MessageToPromptConverter.java @@ -25,7 +25,7 @@ import org.springframework.ai.chat.messages.MessageType; * Converts a list of messages to a prompt for bedrock models. * * @author Christian Tzolov - * @since 0.8.0 + * @since 1.0.0 */ public class MessageToPromptConverter { @@ -75,10 +75,8 @@ public class MessageToPromptConverter { .collect(Collectors.joining(System.lineSeparator())); // Related to: https://github.com/spring-projects/spring-ai/issues/404 - final String prompt = systemMessages + this.lineSeparator + this.lineSeparator + userMessages - + this.lineSeparator + ASSISTANT_PROMPT; - - return prompt; + return systemMessages + this.lineSeparator + this.lineSeparator + userMessages + this.lineSeparator + + ASSISTANT_PROMPT; } protected String messageToString(Message message) { @@ -89,7 +87,7 @@ public class MessageToPromptConverter { return humanPrompt + " " + message.getContent(); case ASSISTANT: return assistantPrompt + " " + message.getContent(); - case FUNCTION: + case TOOL: throw new IllegalArgumentException("Tool execution results are not supported for Bedrock models"); } 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 39deeb635..791f76cf3 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,18 +15,12 @@ */ package org.springframework.ai.openai; -import java.util.ArrayList; -import java.util.Base64; -import java.util.HashMap; -import java.util.HashSet; -import java.util.List; -import java.util.Map; -import java.util.Optional; -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.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.metadata.RateLimit; import org.springframework.ai.chat.model.ChatModel; @@ -36,15 +30,15 @@ import org.springframework.ai.chat.model.StreamingChatModel; 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.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion.Choice; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionFinishReason; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage; +import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.ChatCompletionFunction; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.MediaContent; -import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.Role; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.ToolCall; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest; import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadata; @@ -55,8 +49,16 @@ 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} @@ -77,9 +79,7 @@ import reactor.core.publisher.Flux; * @see StreamingChatModel * @see OpenAiApi */ -public class OpenAiChatModel extends - AbstractFunctionCallSupport> - implements ChatModel { +public class OpenAiChatModel extends AbstractToolCallSupport implements ChatModel { private static final Logger logger = LoggerFactory.getLogger(OpenAiChatModel.class); @@ -145,14 +145,25 @@ public class OpenAiChatModel extends return this.retryTemplate.execute(ctx -> { - ResponseEntity completionEntity = this.callWithFunctionSupport(request); + ResponseEntity completionEntity = this.openAiApi.chatCompletionEntity(request); var chatCompletion = completionEntity.getBody(); + if (chatCompletion == null) { logger.warn("No chat completion returned for prompt: {}", prompt); return new ChatResponse(List.of()); } + if (isToolFunctionCall(chatCompletion)) { + List toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(), + chatCompletion); + // Recursively call the call method with the tool call message + // conversation that contains the call responses. + + return this.call(new Prompt(toolCallMessageConversation, prompt.getOptions())); + } + + // Non function calling. RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity); List choices = chatCompletion.choices(); @@ -162,7 +173,10 @@ public class OpenAiChatModel extends } List generations = choices.stream().map(choice -> { - var generation = new Generation(choice.message().content(), toMap(chatCompletion.id(), choice)); + Map metadata = Map.of("id", chatCompletion.id(), "role", + choice.message().role() != null ? choice.message().role().name() : "", "finishReason", + choice.finishReason() != null ? choice.finishReason().name() : ""); + var generation = new Generation(choice.message().content(), metadata); if (choice.finishReason() != null) { generation = generation .withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), null)); @@ -176,20 +190,6 @@ public class OpenAiChatModel extends }); } - private Map toMap(String id, ChatCompletion.Choice choice) { - Map map = new HashMap<>(); - - var message = choice.message(); - if (message.role() != null) { - map.put("role", message.role().name()); - } - if (choice.finishReason() != null) { - map.put("finishReason", choice.finishReason().name()); - } - map.put("id", id); - return map; - } - @Override public Flux stream(Prompt prompt) { @@ -205,16 +205,23 @@ public class OpenAiChatModel extends // Convert the ChatCompletionChunk into a ChatCompletion to be able to reuse // the function call handling logic. - return completionChunks.map(chunk -> chunkToChatCompletion(chunk)) - .switchMap( - cc -> handleFunctionCallOrReturnStream(request, Flux.just(ResponseEntity.of(Optional.of(cc))))) - .map(ResponseEntity::getBody) - .map(chatCompletion -> { + return completionChunks.map(this::chunkToChatCompletion).switchMap(chatCompletion -> { + + if (this.isToolFunctionCall(chatCompletion)) { + var toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(), + chatCompletion); + // Recursively call the stream method with the tool call message + // conversation that contains the call responses. + return this.stream(new Prompt(toolCallMessageConversation, prompt.getOptions())); + } + + // Non function calling. + return Mono.just(chatCompletion).map(chatCompletion2 -> { try { @SuppressWarnings("null") - String id = chatCompletion.id(); + String id = chatCompletion2.id(); - List generations = chatCompletion.choices().stream().map(choice -> { + List generations = chatCompletion2.choices().stream().map(choice -> { if (choice.message().role() != null) { roleMap.putIfAbsent(id, choice.message().role().name()); } @@ -228,8 +235,8 @@ public class OpenAiChatModel extends return generation; }).toList(); - if (chatCompletion.usage() != null) { - return new ChatResponse(generations, OpenAiChatResponseMetadata.from(chatCompletion)); + if (chatCompletion2.usage() != null) { + return new ChatResponse(generations, OpenAiChatResponseMetadata.from(chatCompletion2)); } else { return new ChatResponse(generations); @@ -241,9 +248,33 @@ public class OpenAiChatModel extends } }); + }); }); } + private List handleToolCallRequests(List previousMessages, ChatCompletion chatCompletion) { + + ChatCompletionMessage nativeAssistantMessage = this.extractAssistantMessage(chatCompletion); + + List assistantToolCalls = nativeAssistantMessage.toolCalls() + .stream() + .map(toolCall -> new AssistantMessage.ToolCall(toolCall.id(), "function", toolCall.function().name(), + toolCall.function().arguments())) + .toList(); + + AssistantMessage assistantMessage = new AssistantMessage(nativeAssistantMessage.content(), Map.of(), + assistantToolCalls); + + List toolResponseMessages = this.executeFuncitons(assistantMessage); + + // History + List messages = new ArrayList<>(previousMessages); + messages.add(assistantMessage); + messages.addAll(toolResponseMessages); + + return messages; + } + /** * Convert the ChatCompletionChunk into a ChatCompletion. The Usage is set to null. * @param chunk the ChatCompletionChunk to convert @@ -252,13 +283,18 @@ public class OpenAiChatModel extends private OpenAiApi.ChatCompletion chunkToChatCompletion(OpenAiApi.ChatCompletionChunk chunk) { List choices = chunk.choices() .stream() - .map(cc -> new Choice(cc.finishReason(), cc.index(), cc.delta(), cc.logprobs())) + .map(chunkChoice -> new Choice(chunkChoice.finishReason(), chunkChoice.index(), chunkChoice.delta(), + chunkChoice.logprobs())) .toList(); return new OpenAiApi.ChatCompletion(chunk.id(), choices, chunk.created(), chunk.model(), chunk.systemFingerprint(), "chat.completion", chunk.usage()); } + private ChatCompletionMessage extractAssistantMessage(ChatCompletion chatCompletion) { + return chatCompletion.choices().iterator().next().message(); + } + /** * Accessible for testing. */ @@ -266,24 +302,47 @@ public class OpenAiChatModel extends Set functionsForThisRequest = new HashSet<>(); - List chatCompletionMessages = prompt.getInstructions().stream().map(m -> { - Object content; - if (CollectionUtils.isEmpty(m.getMedia())) { - content = m.getContent(); + List chatCompletionMessages = prompt.getInstructions().stream().map(message -> { + if (message.getMessageType() == MessageType.USER || message.getMessageType() == MessageType.SYSTEM) { + Object content; + if (CollectionUtils.isEmpty(message.getMedia())) { + content = message.getContent(); + } + else { + List contentList = new ArrayList<>(List.of(new MediaContent(message.getContent()))); + + contentList.addAll(message.getMedia() + .stream() + .map(media -> new MediaContent( + new MediaContent.ImageUrl(this.fromMediaData(media.getMimeType(), media.getData())))) + .toList()); + + content = contentList; + } + + return new ChatCompletionMessage(content, + ChatCompletionMessage.Role.valueOf(message.getMessageType().name())); + } + else if (message.getMessageType() == MessageType.ASSISTANT) { + var assistantMessage = (AssistantMessage) message; + List toolCalls = null; + if (!CollectionUtils.isEmpty(assistantMessage.getToolCalls())) { + toolCalls = assistantMessage.getToolCalls().stream().map(toolCall -> { + var function = new ChatCompletionFunction(toolCall.name(), toolCall.arguments()); + return new ToolCall(toolCall.id(), toolCall.type(), function); + }).toList(); + } + return new ChatCompletionMessage(assistantMessage.getContent(), ChatCompletionMessage.Role.ASSISTANT, + null, null, toolCalls); + } + else if (message.getMessageType() == MessageType.TOOL) { + ToolResponseMessage toolMessage = (ToolResponseMessage) message; + return new ChatCompletionMessage(toolMessage.getContent(), ChatCompletionMessage.Role.TOOL, + toolMessage.getName(), toolMessage.getId(), null); } else { - List contentList = new ArrayList<>(List.of(new MediaContent(m.getContent()))); - - contentList.addAll(m.getMedia() - .stream() - .map(media -> new MediaContent( - new MediaContent.ImageUrl(this.fromMediaData(media.getMimeType(), media.getData())))) - .toList()); - - content = contentList; + throw new IllegalArgumentException("Unsupported message type: " + message.getMessageType()); } - - return new ChatCompletionMessage(content, ChatCompletionMessage.Role.valueOf(m.getMessageType().name())); }).toList(); ChatCompletionRequest request = new ChatCompletionRequest(chatCompletionMessages, stream); @@ -351,66 +410,12 @@ public class OpenAiChatModel extends } @Override - protected ChatCompletionRequest doCreateToolResponseRequest(ChatCompletionRequest previousRequest, - ChatCompletionMessage responseMessage, List conversationHistory) { - - // Every tool-call item requires a separate function call and a response (TOOL) - // message. - for (ToolCall toolCall : responseMessage.toolCalls()) { - - var functionName = toolCall.function().name(); - String functionArguments = toolCall.function().arguments(); - - if (!this.functionCallbackRegister.containsKey(functionName)) { - throw new IllegalStateException("No function callback found for function name: " + functionName); - } - - String functionResponse = this.functionCallbackRegister.get(functionName).call(functionArguments); - - // Add the function response to the conversation. - conversationHistory - .add(new ChatCompletionMessage(functionResponse, Role.TOOL, functionName, toolCall.id(), null)); - } - - // Recursively call chatCompletionWithTools until the model doesn't call a - // functions anymore. - ChatCompletionRequest newRequest = new ChatCompletionRequest(conversationHistory, previousRequest.stream()); - newRequest = ModelOptionsUtils.merge(newRequest, previousRequest, ChatCompletionRequest.class); - - return newRequest; - } - - @Override - protected List doGetUserMessages(ChatCompletionRequest request) { - return request.messages(); - } - - @Override - protected ChatCompletionMessage doGetToolResponseMessage(ResponseEntity chatCompletion) { - return chatCompletion.getBody().choices().iterator().next().message(); - } - - @Override - protected ResponseEntity doChatCompletion(ChatCompletionRequest request) { - return this.openAiApi.chatCompletionEntity(request); - } - - @Override - protected Flux> doChatCompletionStream(ChatCompletionRequest request) { - return this.openAiApi.chatCompletionStream(request) - .map(this::chunkToChatCompletion) - .map(Optional::ofNullable) - .map(ResponseEntity::of); - } - - @Override - protected boolean isToolFunctionCall(ResponseEntity chatCompletion) { - var body = chatCompletion.getBody(); - if (body == null) { + protected boolean isToolFunctionCall(ChatCompletion chatCompletion) { + if (chatCompletion == null) { return false; } - var choices = body.choices(); + var choices = chatCompletion.choices(); if (CollectionUtils.isEmpty(choices)) { return false; } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java index 0fc7d9582..40e61bdce 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java @@ -233,9 +233,7 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions { } public Builder withStreamUsage(boolean enableStreamUsage) { - if (enableStreamUsage) { - this.options.streamOptions = (enableStreamUsage) ? StreamOptions.INCLUDE_USAGE : null; - } + this.options.streamOptions = (enableStreamUsage) ? StreamOptions.INCLUDE_USAGE : null; return this; } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModel3IT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModel3IT.java new file mode 100644 index 000000000..fdf846431 --- /dev/null +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModel3IT.java @@ -0,0 +1,128 @@ +/* + * 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.openai.chat; + +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.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +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.Prompt; +import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.openai.OpenAiChatModel; +import org.springframework.ai.openai.OpenAiChatOptions; +import org.springframework.ai.openai.api.OpenAiApi; +import org.springframework.ai.openai.api.tool.MockWeatherService; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.annotation.Bean; +import reactor.core.publisher.Flux; + +import java.util.ArrayList; +import java.util.List; +import java.util.stream.Collectors; + +import static org.assertj.core.api.Assertions.assertThat; + +@SpringBootTest(classes = OpenAiChatModel3IT.Config.class) +@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") +class OpenAiChatModel3IT { + + private static final Logger logger = LoggerFactory.getLogger(OpenAiChatModel3IT.class); + + @Autowired + ChatModel chatModel; + + @Test + void functionCallTest() { + + UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + + List messages = new ArrayList<>(List.of(userMessage)); + + var promptOptions = OpenAiChatOptions.builder() + .withModel(OpenAiApi.ChatModel.GPT_4_O.getValue()) + .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) + .withName("getCurrentWeather") + .withDescription("Get the weather in location") + .withResponseConverter((response) -> "" + response.temp() + response.unit()) + .build())) + .build(); + + ChatResponse response = chatModel.call(new Prompt(messages, promptOptions)); + + 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"); + } + + @Test + void streamFunctionCallTest() { + + UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + + List messages = new ArrayList<>(List.of(userMessage)); + + var promptOptions = OpenAiChatOptions.builder() + // .withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue()) + .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) + .withName("getCurrentWeather") + .withDescription("Get the weather in location") + .withResponseConverter((response) -> "" + response.temp() + response.unit()) + .build())) + .build(); + + Flux response = chatModel.stream(new Prompt(messages, promptOptions)); + + String content = response.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .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"); + } + + @SpringBootConfiguration + static class Config { + + @Bean + public OpenAiApi chatCompletionApi() { + return new OpenAiApi(System.getenv("OPENAI_API_KEY")); + } + + @Bean + public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { + return new OpenAiChatModel(openAiApi); + } + + } + +} \ No newline at end of file diff --git a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/utils/MessageToPromptConverter.java b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/utils/MessageToPromptConverter.java index 227fd0fc3..75be3e173 100644 --- a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/utils/MessageToPromptConverter.java +++ b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/utils/MessageToPromptConverter.java @@ -71,7 +71,7 @@ public class MessageToPromptConverter { return humanPrompt + message.getContent(); case ASSISTANT: return assistantPrompt + message.getContent(); - case FUNCTION: + case TOOL: throw new IllegalArgumentException(TOOL_EXECUTION_NOT_SUPPORTED_FOR_WAI_MODELS); } diff --git a/models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/utils/MessageToPromptConverterTest.java b/models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/utils/MessageToPromptConverterTest.java index a1dcaa608..5d22477c6 100644 --- a/models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/utils/MessageToPromptConverterTest.java +++ b/models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/utils/MessageToPromptConverterTest.java @@ -23,7 +23,6 @@ import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.UserMessage; -import org.springframework.ai.chat.messages.FunctionMessage; import java.util.List; @@ -54,14 +53,6 @@ public class MessageToPromptConverterTest { Assert.assertEquals(expected, converter.messageToString(assistantMessage)); } - @Disabled - public void testFunctionMessageType() { - Message functionMessage = new FunctionMessage("Function message"); - Exception exception = Assert.assertThrows(IllegalArgumentException.class, () -> { - converter.messageToString(functionMessage); - }); - } - @Test public void testSystemMessageType() { Message systemMessage = new SystemMessage("System message"); @@ -105,12 +96,4 @@ public class MessageToPromptConverterTest { Assert.assertEquals(expected, converter.toPrompt(messages)); } - @Disabled - public void testUnsupportedMessageType() { - List messages = List.of(new FunctionMessage("Unsupported message")); - Exception exception = Assert.assertThrows(IllegalArgumentException.class, () -> { - converter.toPrompt(messages); - }); - } - } diff --git a/qodana.yaml b/qodana.yaml new file mode 100644 index 000000000..08a316bd7 --- /dev/null +++ b/qodana.yaml @@ -0,0 +1,31 @@ +#-------------------------------------------------------------------------------# +# Qodana analysis is configured by qodana.yaml file # +# https://www.jetbrains.com/help/qodana/qodana-yaml.html # +#-------------------------------------------------------------------------------# +version: "1.0" + +#Specify inspection profile for code analysis +profile: + name: qodana.starter + +#Enable inspections +#include: +# - name: + +#Disable inspections +#exclude: +# - name: +# paths: +# - + +projectJDK: 17 #(Applied in CI/CD pipeline) + +#Execute shell command before Qodana execution (Applied in CI/CD pipeline) +#bootstrap: sh ./prepare-qodana.sh + +#Install IDE plugins before Qodana execution (Applied in CI/CD pipeline) +#plugins: +# - id: #(plugin id can be found at https://plugins.jetbrains.com) + +#Specify Qodana linter for analysis (Applied in CI/CD pipeline) +linter: jetbrains/qodana-jvm:latest diff --git a/spring-ai-core/src/main/java/org/springframework/ai/aot/SpringAiCoreRuntimeHints.java b/spring-ai-core/src/main/java/org/springframework/ai/aot/SpringAiCoreRuntimeHints.java index 9fbc7c16c..393687bec 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/aot/SpringAiCoreRuntimeHints.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/aot/SpringAiCoreRuntimeHints.java @@ -18,7 +18,7 @@ package org.springframework.ai.aot; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.AbstractMessage; import org.springframework.ai.chat.messages.AssistantMessage; -import org.springframework.ai.chat.messages.FunctionMessage; +import org.springframework.ai.chat.messages.ToolResponseMessage; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.UserMessage; @@ -41,7 +41,7 @@ public class SpringAiCoreRuntimeHints implements RuntimeHintsRegistrar { @Override public void registerHints(@NonNull RuntimeHints hints, @Nullable ClassLoader classLoader) { - var chatTypes = Set.of(AbstractMessage.class, AssistantMessage.class, FunctionMessage.class, Message.class, + var chatTypes = Set.of(AbstractMessage.class, AssistantMessage.class, ToolResponseMessage.class, Message.class, MessageType.class, UserMessage.class, SystemMessage.class, FunctionCallbackContext.class, FunctionCallback.class, FunctionCallbackWrapper.class); for (var c : chatTypes) { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AssistantMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AssistantMessage.java index f8b890416..8791ab9f1 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AssistantMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AssistantMessage.java @@ -15,28 +15,64 @@ */ package org.springframework.ai.chat.messages; +import org.springframework.util.Assert; + +import java.util.List; import java.util.Map; +import java.util.Objects; /** * Lets the generative know the content was generated as a response to the user. This role * indicates messages that the generative has previously generated in the conversation. By * including assistant messages in the series, you provide context to the generative about * prior exchanges in the conversation. + * + * @author Mark Pollack + * @author Christian Tzolov + * @since 1.0.0 */ public class AssistantMessage extends AbstractMessage { + public record ToolCall(String id, String type, String name, String arguments) { + } + + private final List toolCalls; + public AssistantMessage(String content) { - super(MessageType.ASSISTANT, content); + this(content, Map.of()); } public AssistantMessage(String content, Map properties) { + this(content, properties, List.of()); + } + + public AssistantMessage(String content, Map properties, List toolCalls) { super(MessageType.ASSISTANT, content, properties); + Assert.notNull(toolCalls, "Tool calls must not be null"); + this.toolCalls = toolCalls; + } + + public List getToolCalls() { + return this.toolCalls; } @Override - public String toString() { - return "AssistantMessage{" + "content='" + getContent() + '\'' + ", properties=" + metadata + ", messageType=" - + messageType + '}'; + public int hashCode() { + return Objects.hash(this.messageType, this.getContent(), this.metadata, this.toolCalls); + } + + @Override + public boolean equals(Object obj) { + if (this == obj) { + return true; + } + if (obj == null || getClass() != obj.getClass()) { + return false; + } + AssistantMessage other = (AssistantMessage) obj; + return this.messageType == other.messageType && Objects.equals(this.getContent(), other.getContent()) + && Objects.equals(this.getMetadata(), other.getMetadata()) + && Objects.equals(this.getToolCalls(), other.getToolCalls()); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/FunctionMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/FunctionMessage.java deleted file mode 100644 index a05ef5226..000000000 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/FunctionMessage.java +++ /dev/null @@ -1,40 +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.chat.messages; - -import java.util.Map; - -/** - * The FunctionMessage class represents a message with a function content in a chat - * application. - */ -public class FunctionMessage extends AbstractMessage { - - public FunctionMessage(String content) { - super(MessageType.FUNCTION, content); - } - - public FunctionMessage(String content, Map properties) { - super(MessageType.FUNCTION, content, properties); - } - - @Override - public String toString() { - return "FunctionMessage{" + "content='" + getContent() + '\'' + ", properties=" + metadata + ", messageType=" - + messageType + '}'; - } - -} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/MessageType.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/MessageType.java index c1c88dd43..7603ab39f 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/MessageType.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/MessageType.java @@ -46,9 +46,9 @@ public enum MessageType { /** * A message of the type 'function' passed as input Messages with a function content * in a chat application. - * @see FunctionMessage + * @see ToolResponseMessage */ - FUNCTION("function"); + TOOL("tool"); private final String value; 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 new file mode 100644 index 000000000..66af01488 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/ToolResponseMessage.java @@ -0,0 +1,77 @@ +/* + * 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.chat.messages; + +import java.util.Map; +import java.util.Objects; + +/** + * The FunctionMessage class represents a message with a function content in a chat + * application. + * + * @author Christian Tzolov + * @since 1.0.0 + */ +public class ToolResponseMessage extends AbstractMessage { + + private final String id; + + private final String name; + + public ToolResponseMessage(String id, String name, String content) { + this(id, name, content, Map.of()); + } + + public ToolResponseMessage(String id, String name, String content, Map metadata) { + super(MessageType.TOOL, content, metadata); + this.id = id; + this.name = name; + } + + public String getId() { + return id; + } + + public String getName() { + return name; + } + + @Override + public int hashCode() { + return Objects.hash(this.id, this.name, getContent(), this.metadata, this.messageType); + } + + @Override + public boolean equals(Object obj) { + if (this == obj) { + return true; + } + if (obj == null || getClass() != obj.getClass()) { + 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; + } + + @Override + public String toString() { + return "FunctionMessage [id=" + id + ", name=" + name + ", messageType=" + messageType + ", textContent=" + + textContent + "]"; + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/Prompt.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/Prompt.java index 531564f0a..0b2d2e964 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/Prompt.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/Prompt.java @@ -21,7 +21,7 @@ import java.util.List; import java.util.Objects; import org.springframework.ai.chat.messages.AssistantMessage; -import org.springframework.ai.chat.messages.FunctionMessage; +import org.springframework.ai.chat.messages.ToolResponseMessage; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.UserMessage; @@ -115,11 +115,13 @@ public class Prompt implements ModelRequest> { else if (message instanceof SystemMessage) { messagesCopy.add(new SystemMessage(message.getContent())); } - else if (message instanceof AssistantMessage) { - messagesCopy.add(new AssistantMessage(message.getContent(), message.getMetadata())); + else if (message instanceof AssistantMessage assistantMessage) { + messagesCopy.add(new AssistantMessage(assistantMessage.getContent(), assistantMessage.getMetadata(), + assistantMessage.getToolCalls())); } - else if (message instanceof FunctionMessage) { - messagesCopy.add(new FunctionMessage(message.getContent(), message.getMetadata())); + else if (message instanceof ToolResponseMessage toolResponseMessage) { + messagesCopy.add(new ToolResponseMessage(toolResponseMessage.getId(), toolResponseMessage.getName(), + toolResponseMessage.getContent(), toolResponseMessage.getMetadata())); } else { throw new IllegalArgumentException("Unsupported message type: " + message.getClass().getName()); 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 new file mode 100644 index 000000000..c5f7ad9fe --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java @@ -0,0 +1,154 @@ +/* + * 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.model.function; + +import java.util.ArrayList; +import java.util.HashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.concurrent.ConcurrentHashMap; + +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.ToolResponseMessage; +import org.springframework.util.CollectionUtils; + +/** + * Abstract base class for tool call support. Provides functionality for handling function + * callbacks and executing functions. + * + * @param The response type of the tool call. + * @author Christian Tzolov + * @author Grogdunn + * @since 1.0.0 + */ +public abstract class AbstractToolCallSupport { + + protected final static boolean IS_RUNTIME_CALL = true; + + /** + * The function callback register is used to resolve the function callbacks by name. + */ + protected final Map functionCallbackRegister = new ConcurrentHashMap<>(); + + /** + * The function callback context is used to resolve the function callbacks by name + * from the Spring context. It is optional and usually used with Spring + * auto-configuration. + */ + protected final FunctionCallbackContext functionCallbackContext; + + protected AbstractToolCallSupport(FunctionCallbackContext functionCallbackContext) { + this.functionCallbackContext = functionCallbackContext; + } + + public Map getFunctionCallbackRegister() { + return this.functionCallbackRegister; + } + + protected Set handleFunctionCallbackConfigurations(FunctionCallingOptions options, boolean isRuntimeCall) { + + Set functionToCall = new HashSet<>(); + + if (options != null) { + if (!CollectionUtils.isEmpty(options.getFunctionCallbacks())) { + options.getFunctionCallbacks().stream().forEach(functionCallback -> { + + // Register the tool callback. + if (isRuntimeCall) { + this.functionCallbackRegister.put(functionCallback.getName(), functionCallback); + } + else { + this.functionCallbackRegister.putIfAbsent(functionCallback.getName(), functionCallback); + } + + // Automatically enable the function, usually from prompt callback. + if (isRuntimeCall) { + functionToCall.add(functionCallback.getName()); + } + }); + } + + // Add the explicitly enabled functions. + if (!CollectionUtils.isEmpty(options.getFunctions())) { + functionToCall.addAll(options.getFunctions()); + } + } + + return functionToCall; + } + + /** + * Resolve the function callbacks by name. Retrieve them from the registry or try to + * resolve them from the Application Context. + * @param functionNames Name of function callbacks to retrieve. + * @return list of resolved FunctionCallbacks. + */ + protected List resolveFunctionCallbacks(Set functionNames) { + + List retrievedFunctionCallbacks = new ArrayList<>(); + + for (String functionName : functionNames) { + if (!this.functionCallbackRegister.containsKey(functionName)) { + + if (this.functionCallbackContext != null) { + FunctionCallback functionCallback = this.functionCallbackContext.getFunctionCallback(functionName, + null); + if (functionCallback != null) { + this.functionCallbackRegister.put(functionName, functionCallback); + } + else { + throw new IllegalStateException( + "No function callback [" + functionName + "] fund in tht FunctionCallbackContext"); + } + } + else { + throw new IllegalStateException("No function callback found for name: " + functionName); + } + } + FunctionCallback functionCallback = this.functionCallbackRegister.get(functionName); + + retrievedFunctionCallbacks.add(functionCallback); + } + + return retrievedFunctionCallbacks; + } + + protected List executeFuncitons(AssistantMessage assistantMessage) { + + List toolResponseMessages = new ArrayList<>(); + + for (AssistantMessage.ToolCall toolCall : assistantMessage.getToolCalls()) { + + var functionName = toolCall.name(); + String functionArguments = toolCall.arguments(); + + if (!this.functionCallbackRegister.containsKey(functionName)) { + throw new IllegalStateException("No function callback found for function name: " + functionName); + } + + String functionResponse = this.functionCallbackRegister.get(functionName).call(functionArguments); + + toolResponseMessages.add(new ToolResponseMessage(toolCall.id(), functionName, functionResponse)); + } + + return toolResponseMessages; + + } + + abstract protected boolean isToolFunctionCall(TRes response); + +}