From a7eb28ac17debbaaaf227802eb0385f3daea4fd8 Mon Sep 17 00:00:00 2001 From: Lorenzo Caenazzo Date: Thu, 21 Mar 2024 09:43:29 +0100 Subject: [PATCH] Add real Function Calling Streaming support - Add Java reflection merge utilities that can access Azure private constructors and fields. - Azure merging, creation of flux windows. - Function call grouping for function processing. - Do not perform greedy operation on Flux. - Use "real" streaming on all client on function response. - Gerimi: fix missing method impl. - Mistral AI, OpenAI: fix missing stream flag in doCreateToolResponseRequest. - Fix code formatting. No wildcard imports. - Add Grogdunn to the javadoc authors. - Anthropic 3 API does not support streaming funciton calling yet. --- .../ai/anthropic/AnthropicChatClient.java | 7 + .../ai/anthropic/api/AnthropicApi.java | 9 +- .../ai/anthropic/AnthropicChatClientIT.java | 4 +- .../azure/openai/AzureOpenAiChatClient.java | 71 +++- .../ai/azure/openai/MergeUtils.java | 323 ++++++++++++++++++ .../AzureOpenAiChatClientFunctionCallIT.java | 57 +++- .../openai/function/MockWeatherService.java | 6 +- .../ai/mistralai/MistralAiChatClient.java | 69 ++-- .../ai/openai/OpenAiChatClient.java | 89 ++--- .../gemini/VertexAiGeminiChatClient.java | 52 +-- .../function/AbstractFunctionCallSupport.java | 39 ++- 11 files changed, 605 insertions(+), 121 deletions(-) create mode 100644 models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/MergeUtils.java diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatClient.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatClient.java index 73fddf179..0f9bcbf44 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatClient.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatClient.java @@ -450,4 +450,11 @@ public class AnthropicChatClient extends return response.getBody().content().stream().anyMatch(content -> content.type() == MediaContent.Type.TOOL_USE); } + @Override + protected Flux> doChatCompletionStream(ChatCompletionRequest request) { + // https://docs.anthropic.com/en/docs/tool-use + throw new UnsupportedOperationException( + "Streaming (stream=true) is not yet supported. We plan to add streaming support in a future beta version."); + } + } diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/api/AnthropicApi.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/api/AnthropicApi.java index 46e6119d6..eb5a96289 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/api/AnthropicApi.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/api/AnthropicApi.java @@ -29,6 +29,7 @@ import reactor.core.publisher.Mono; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.retry.RetryUtils; import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatusCode; import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; import org.springframework.util.Assert; @@ -100,7 +101,13 @@ public class AnthropicApi { .defaultStatusHandler(responseErrorHandler) .build(); - this.webClient = WebClient.builder().baseUrl(baseUrl).defaultHeaders(jsonContentHeaders).build(); + this.webClient = WebClient.builder() + .baseUrl(baseUrl) + .defaultHeaders(jsonContentHeaders) + .defaultStatusHandler(HttpStatusCode::isError, + resp -> Mono.just(new RuntimeException("Response exception, Status: [" + resp.statusCode() + + "], Body:[" + resp.bodyToMono(java.lang.String.class) + "]"))) + .build(); } /** diff --git a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java index 7bc7ef2af..92ba0a6eb 100644 --- a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java +++ b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java @@ -205,7 +205,7 @@ class AnthropicChatClientIT { .withModel(AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue()) .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) .withName("getCurrentWeather") - .withDescription("Get the weather in location") + .withDescription("Get the weather in location. Return temperature in 36°F or 36°C format.") .build())) .build(); @@ -213,7 +213,7 @@ class AnthropicChatClientIT { logger.info("Response: {}", response); - Generation generation = response.getResults().get(0); + Generation generation = response.getResult(); assertThat(generation.getOutput().getContent()).containsAnyOf("30.0", "30"); assertThat(generation.getOutput().getContent()).containsAnyOf("10.0", "10"); assertThat(generation.getOutput().getContent()).containsAnyOf("15.0", "15"); diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatClient.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatClient.java index 236a902d8..a49a42ff5 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatClient.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatClient.java @@ -15,11 +15,6 @@ */ package org.springframework.ai.azure.openai; -import java.util.Collections; -import java.util.HashSet; -import java.util.List; -import java.util.Set; - import com.azure.ai.openai.OpenAIClient; import com.azure.ai.openai.models.ChatChoice; import com.azure.ai.openai.models.ChatCompletions; @@ -33,15 +28,14 @@ import com.azure.ai.openai.models.ChatRequestMessage; import com.azure.ai.openai.models.ChatRequestSystemMessage; import com.azure.ai.openai.models.ChatRequestToolMessage; import com.azure.ai.openai.models.ChatRequestUserMessage; -import com.azure.ai.openai.models.ChatResponseMessage; import com.azure.ai.openai.models.CompletionsFinishReason; import com.azure.ai.openai.models.ContentFilterResultsForPrompt; +import com.azure.ai.openai.models.FunctionCall; import com.azure.ai.openai.models.FunctionDefinition; import com.azure.core.util.BinaryData; import com.azure.core.util.IterableStream; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import reactor.core.publisher.Flux; import org.springframework.ai.azure.openai.metadata.AzureOpenAiChatResponseMetadata; import org.springframework.ai.chat.ChatClient; @@ -59,6 +53,14 @@ import org.springframework.ai.model.function.AbstractFunctionCallSupport; import org.springframework.ai.model.function.FunctionCallbackContext; import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; +import reactor.core.publisher.Flux; + +import java.util.Collections; +import java.util.HashSet; +import java.util.List; +import java.util.Optional; +import java.util.Set; +import java.util.concurrent.atomic.AtomicBoolean; /** * {@link ChatClient} implementation for {@literal Microsoft Azure AI} backed by @@ -68,6 +70,7 @@ import org.springframework.util.CollectionUtils; * @author Ueibin Kim * @author John Blum * @author Christian Tzolov + * @author Grogdunn * @see ChatClient * @see com.azure.ai.openai.OpenAIClient */ @@ -158,17 +161,42 @@ public class AzureOpenAiChatClient IterableStream chatCompletionsStream = this.openAIClient .getChatCompletionsStream(options.getModel(), options); - return Flux.fromStream(chatCompletionsStream.stream() + Flux chatCompletionsFlux = Flux.fromIterable(chatCompletionsStream); + + final var isFunctionCall = new AtomicBoolean(false); + final var accessibleChatCompletionsFlux = chatCompletionsFlux // Note: the first chat completions can be ignored when using Azure OpenAI // service which is a known service bug. .skip(1) - .map(ChatCompletions::getChoices) - .flatMap(List::stream) + .map(chatCompletions -> { + final var toolCalls = chatCompletions.getChoices().get(0).getDelta().getToolCalls(); + isFunctionCall.set(toolCalls != null && !toolCalls.isEmpty()); + return chatCompletions; + }) + .windowUntil(chatCompletions -> { + if (isFunctionCall.get() && chatCompletions.getChoices() + .get(0) + .getFinishReason() == CompletionsFinishReason.TOOL_CALLS) { + isFunctionCall.set(false); + return true; + } + return false; + }, false) + .concatMapIterable(window -> { + final var reduce = window.reduce(MergeUtils.emptyChatCompletions(), MergeUtils::mergeChatCompletions); + return List.of(reduce); + }) + .flatMap(mono -> mono); + return accessibleChatCompletionsFlux + .switchMap(accessibleChatCompletions -> handleFunctionCallOrReturnStream(options, + Flux.just(accessibleChatCompletions))) + .flatMapIterable(ChatCompletions::getChoices) .map(choice -> { - var content = (choice.getDelta() != null) ? choice.getDelta().getContent() : null; + var content = Optional.ofNullable(choice.getMessage()).orElse(choice.getDelta()).getContent(); var generation = new Generation(content).withGenerationMetadata(generateChoiceMetadata(choice)); return new ChatResponse(List.of(generation)); - })); + }); + } /** @@ -522,9 +550,17 @@ public class AzureOpenAiChatClient @Override protected ChatRequestMessage doGetToolResponseMessage(ChatCompletions response) { - ChatResponseMessage responseMessage = response.getChoices().get(0).getMessage(); + final var accessibleChatChoice = response.getChoices().get(0); + var responseMessage = Optional.ofNullable(accessibleChatChoice.getMessage()) + .orElse(accessibleChatChoice.getDelta()); ChatRequestAssistantMessage assistantMessage = new ChatRequestAssistantMessage(""); - assistantMessage.setToolCalls(responseMessage.getToolCalls()); + final var toolCalls = responseMessage.getToolCalls(); + assistantMessage.setToolCalls(toolCalls.stream().map(tc -> { + final var tc1 = (ChatCompletionsFunctionToolCall) tc; + var toDowncast = new ChatCompletionsFunctionToolCall(tc.getId(), + new FunctionCall(tc1.getFunction().getName(), tc1.getFunction().getArguments())); + return ((ChatCompletionsToolCall) toDowncast); + }).toList()); return assistantMessage; } @@ -533,6 +569,11 @@ public class AzureOpenAiChatClient return this.openAIClient.getChatCompletions(request.getModel(), request); } + @Override + protected Flux doChatCompletionStream(ChatCompletionsOptions request) { + return Flux.fromIterable(this.openAIClient.getChatCompletionsStream(request.getModel(), request)); + } + @Override protected boolean isToolFunctionCall(ChatCompletions chatCompletions) { @@ -549,4 +590,4 @@ public class AzureOpenAiChatClient return choice.getFinishReason() == CompletionsFinishReason.TOOL_CALLS; } -} \ No newline at end of file +} diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/MergeUtils.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/MergeUtils.java new file mode 100644 index 000000000..a4e995937 --- /dev/null +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/MergeUtils.java @@ -0,0 +1,323 @@ +/* + * 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.azure.openai; + +import java.lang.reflect.Constructor; +import java.lang.reflect.Field; +import java.time.OffsetDateTime; +import java.util.ArrayList; +import java.util.List; +import java.util.Objects; + +import com.azure.ai.openai.models.AzureChatExtensionsMessageContext; +import com.azure.ai.openai.models.ChatChoice; +import com.azure.ai.openai.models.ChatCompletions; +import com.azure.ai.openai.models.ChatCompletionsFunctionToolCall; +import com.azure.ai.openai.models.ChatCompletionsToolCall; +import com.azure.ai.openai.models.ChatResponseMessage; +import com.azure.ai.openai.models.CompletionsFinishReason; +import com.azure.ai.openai.models.CompletionsUsage; +import com.azure.ai.openai.models.ContentFilterResultsForChoice; +import com.azure.ai.openai.models.ContentFilterResultsForPrompt; +import com.azure.ai.openai.models.FunctionCall; + +import org.springframework.util.Assert; +import org.springframework.util.CollectionUtils; + +/** + * Utility class for merging ChatCompletions instances and their associated objects. Uses + * reflection to create instances with private constructors and set private fields. + * + * @author Grogdunn + * @author Christian Tzolov + * @since 1.0.0 + */ +public class MergeUtils { + + /** + * Create a new instance of the given class. Can be used to create instances with + * private constructors. + * @param the type of the class to be created. + * @param clazz the class to create an instance of. + * @param args the arguments to pass to the constructor. + * @return a new instance of the given class. + */ + private static T newInstance(Class clazz, Object... args) { + return newInstance(0, clazz, args); + } + + /** + * Create a new instance of the given class using the constructor at the given index. + * Can be used to create instances with private constructors. + * @param the type of the class to be created. + * @param index the index of the constructor to use. + * @param clazz the class to create an instance of. + * @param args the arguments to pass to the constructor. + * @return a new instance of the given class. + */ + private static T newInstance(int index, Class clazz, Object... args) { + try { + @SuppressWarnings("unchecked") + Constructor constructor = (Constructor) clazz.getDeclaredConstructors()[index]; + constructor.setAccessible(true); + return constructor.newInstance(args); + } + catch (Exception e) { + throw new RuntimeException(e); + } + } + + /** + * Set the value of a private field in the given class instance. + * @param classInstance the class instance to set the field on. + * @param fieldName the name of the field to set. + * @param fieldValue the value to set the field to. + */ + private static void setField(Object classInstance, String fieldName, Object fieldValue) { + try { + Field field = classInstance.getClass().getDeclaredField(fieldName); + field.setAccessible(true); + field.set(classInstance, fieldValue); + } + catch (Exception e) { + throw new RuntimeException(e); + } + } + + /** + * @return an empty ChatCompletions instance. + */ + public static ChatCompletions emptyChatCompletions() { + String id = null; + List choices = new ArrayList<>(); + CompletionsUsage usage = null; + long createdAt = 0; + ChatCompletions chatCompletionsInstance = newInstance(ChatCompletions.class, id, createdAt, choices, usage); + List promptFilterResults = new ArrayList<>(); + setField(chatCompletionsInstance, "promptFilterResults", promptFilterResults); + String systemFingerprint = null; + setField(chatCompletionsInstance, "systemFingerprint", systemFingerprint); + + return chatCompletionsInstance; + } + + /** + * Merge two ChatCompletions instances into a single ChatCompletions instance. + * @param left the left ChatCompletions instance. + * @param right the right ChatCompletions instance. + * @return a merged ChatCompletions instance. + */ + public static ChatCompletions mergeChatCompletions(ChatCompletions left, ChatCompletions right) { + + Assert.isTrue(left != null, ""); + if (right == null) { + Assert.isTrue(left.getId() != null, ""); + return left; + } + Assert.isTrue(left.getId() != null || right.getId() != null, ""); + + String id = left.getId() != null ? left.getId() : right.getId(); + + List choices = null; + if (right.getChoices() == null) { + choices = left.getChoices(); + } + else { + if (CollectionUtils.isEmpty(left.getChoices())) { + choices = right.getChoices(); + } + else { + choices = List.of(mergeChatChoice(left.getChoices().get(0), right.getChoices().get(0))); + } + } + + // For these properties if right contains that use it! + CompletionsUsage usage = right.getUsage() == null ? left.getUsage() : right.getUsage(); + + OffsetDateTime createdAt = left.getCreatedAt().isAfter(right.getCreatedAt()) ? left.getCreatedAt() + : right.getCreatedAt(); + + ChatCompletions instance = newInstance(1, ChatCompletions.class, id, createdAt, choices, usage); + + List promptFilterResults = right.getPromptFilterResults() == null + ? left.getPromptFilterResults() : right.getPromptFilterResults(); + setField(instance, "promptFilterResults", promptFilterResults); + + String systemFingerprint = right.getSystemFingerprint() == null ? left.getSystemFingerprint() + : right.getSystemFingerprint(); + setField(instance, "systemFingerprint", systemFingerprint); + return instance; + } + + /** + * Merge two ChatChoice instances into a single ChatChoice instance. + * @param left the left ChatChoice instance to merge. + * @param right the right ChatChoice instance to merge. + * @return a merged ChatChoice instance. + */ + private static ChatChoice mergeChatChoice(ChatChoice left, ChatChoice right) { + + int index = Math.max(left.getIndex(), right.getIndex()); + + CompletionsFinishReason finishReason = left.getFinishReason() != null ? left.getFinishReason() + : right.getFinishReason(); + + var logprobs = left.getLogprobs() != null ? left.getLogprobs() : right.getLogprobs(); + + final ChatChoice instance = newInstance(ChatChoice.class, logprobs, index, finishReason); + + ChatResponseMessage message = null; + if (left.getMessage() == null) { + message = right.getMessage(); + } + else { + message = mergeChatResponseMessage(left.getMessage(), right.getMessage()); + } + + setField(instance, "message", message); + + ChatResponseMessage delta = null; + if (left.getDelta() == null) { + delta = right.getDelta(); + } + else { + delta = mergeChatResponseMessage(left.getDelta(), right.getDelta()); + } + setField(instance, "delta", delta); + + ContentFilterResultsForChoice contentFilterResults = left.getContentFilterResults() != null + ? left.getContentFilterResults() : right.getContentFilterResults(); + setField(instance, "contentFilterResults", contentFilterResults); + + var finishDetails = left.getFinishDetails() != null ? left.getFinishDetails() : right.getFinishDetails(); + setField(instance, "finishDetails", finishDetails); + + var enhancements = left.getEnhancements() != null ? left.getEnhancements() : right.getEnhancements(); + setField(instance, "enhancements", enhancements); + + return instance; + } + + /** + * Merge two ChatResponseMessage instances into a single ChatResponseMessage instance. + * @param left the left ChatResponseMessage instance to merge. + * @param right the right ChatResponseMessage instance to merge. + * @return a merged ChatResponseMessage instance. + */ + private static ChatResponseMessage mergeChatResponseMessage(ChatResponseMessage left, ChatResponseMessage right) { + + var role = left.getRole() != null ? left.getRole() : right.getRole(); + String content = null; + if (left.getContent() != null && right.getContent() != null) { + content = left.getContent().concat(right.getContent()); + } + else if (left.getContent() == null) { + content = right.getContent(); + } + else { + content = left.getContent(); + } + + ChatResponseMessage instance = newInstance(ChatResponseMessage.class, role, content); + + List toolCalls = new ArrayList<>(); + if (left.getToolCalls() == null) { + if (right.getToolCalls() != null) { + toolCalls.addAll(right.getToolCalls()); + } + } + else if (right.getToolCalls() == null) { + toolCalls.addAll(left.getToolCalls()); + } + else { + toolCalls.addAll(left.getToolCalls()); + final var lastToolIndex = toolCalls.size() - 1; + ChatCompletionsToolCall lastTool = toolCalls.get(lastToolIndex); + if (right.getToolCalls().get(0).getId() == null) { + + lastTool = mergeChatCompletionsToolCall(lastTool, right.getToolCalls().get(0)); + + toolCalls.remove(lastToolIndex); + toolCalls.add(lastTool); + } + else { + toolCalls.add(right.getToolCalls().get(0)); + } + } + + setField(instance, "toolCalls", toolCalls); + + FunctionCall functionCall = null; + + if (left.getFunctionCall() == null) { + functionCall = right.getFunctionCall(); + } + else { + functionCall = MergeUtils.mergeFunctionCall(left.getFunctionCall(), right.getFunctionCall()); + } + + setField(instance, "functionCall", functionCall); + + AzureChatExtensionsMessageContext context = left.getContext() != null ? left.getContext() : right.getContext(); + setField(instance, "context", context); + + return instance; + } + + /** + * Merge two ChatCompletionsToolCall instances into a single ChatCompletionsToolCall + * instance. + * @param left the left ChatCompletionsToolCall instance to merge. + * @param right the right ChatCompletionsToolCall instance to merge. + * @return a merged ChatCompletionsToolCall instance. + */ + private static ChatCompletionsToolCall mergeChatCompletionsToolCall(ChatCompletionsToolCall left, + ChatCompletionsToolCall right) { + Assert.isTrue(Objects.equals(left.getType(), right.getType()), + "Cannot merge different type of AccessibleChatCompletionsToolCall"); + if (!"function".equals(left.getType())) { + throw new UnsupportedOperationException("Only function chat completion tool is supported"); + } + + String id = left.getId() != null ? left.getId() : right.getId(); + var mergedFunction = mergeFunctionCall(((ChatCompletionsFunctionToolCall) left).getFunction(), + ((ChatCompletionsFunctionToolCall) right).getFunction()); + + return new ChatCompletionsFunctionToolCall(id, mergedFunction); + } + + /** + * Merge two FunctionCall instances into a single FunctionCall instance. + * @param left the left, input FunctionCall instance. + * @param right the right, input FunctionCall instance. + * @return a merged FunctionCall instance. + */ + private static FunctionCall mergeFunctionCall(FunctionCall left, FunctionCall right) { + var name = left.getName() != null ? left.getName() : right.getName(); + String arguments = null; + if (left.getArguments() != null && right.getArguments() != null) { + arguments = left.getArguments() + right.getArguments(); + } + else if (left.getArguments() == null) { + arguments = right.getArguments(); + } + else { + arguments = left.getArguments(); + } + return new FunctionCall(name, arguments); + } + +} diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatClientFunctionCallIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatClientFunctionCallIT.java index 3e67dc3c6..08c81ebd1 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatClientFunctionCallIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatClientFunctionCallIT.java @@ -17,6 +17,9 @@ package org.springframework.ai.azure.openai.function; import java.util.ArrayList; import java.util.List; +import java.util.Optional; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.stream.Collectors; import com.azure.ai.openai.OpenAIClient; import com.azure.ai.openai.OpenAIClientBuilder; @@ -29,6 +32,8 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.azure.openai.AzureOpenAiChatClient; import org.springframework.ai.azure.openai.AzureOpenAiChatOptions; import org.springframework.ai.chat.ChatResponse; +import org.springframework.ai.chat.Generation; +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.prompt.Prompt; @@ -37,6 +42,7 @@ 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 static org.assertj.core.api.Assertions.assertThat; @@ -47,6 +53,9 @@ class AzureOpenAiChatClientFunctionCallIT { private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiChatClientFunctionCallIT.class); + @Autowired + private String selectedModel; + @Autowired private AzureOpenAiChatClient chatClient; @@ -58,7 +67,7 @@ class AzureOpenAiChatClientFunctionCallIT { List messages = new ArrayList<>(List.of(userMessage)); var promptOptions = AzureOpenAiChatOptions.builder() - .withDeploymentName("gpt-4-0125-preview") + .withDeploymentName(selectedModel) .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) .withName("getCurrentWeather") .withDescription("Get the current weather in a given location") @@ -75,6 +84,40 @@ class AzureOpenAiChatClientFunctionCallIT { 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 = AzureOpenAiChatOptions.builder() + .withDeploymentName(selectedModel) + .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) + .withName("getCurrentWeather") + .withDescription("Get the current weather in a given location") + .withResponseConverter((response) -> "" + response.temp() + response.unit()) + .build())) + .build(); + + Flux response = chatClient.stream(new Prompt(messages, promptOptions)); + + final var counter = new AtomicInteger(); + String content = response.doOnEach(listSignal -> counter.getAndIncrement()) + .collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + logger.info("Response: {}", content); + assertThat(counter.get()).isGreaterThan(2); + assertThat(content).containsAnyOf("30.0", "30"); + assertThat(content).containsAnyOf("10.0", "10"); + assertThat(content).containsAnyOf("15.0", "15"); + } + @SpringBootConfiguration public static class TestConfiguration { @@ -86,12 +129,14 @@ class AzureOpenAiChatClientFunctionCallIT { } @Bean - public AzureOpenAiChatClient azureOpenAiChatClient(OpenAIClient openAIClient) { + public AzureOpenAiChatClient azureOpenAiChatClient(OpenAIClient openAIClient, String selectedModel) { return new AzureOpenAiChatClient(openAIClient, - AzureOpenAiChatOptions.builder() - .withDeploymentName("gpt-4-0125-preview") - .withMaxTokens(500) - .build()); + AzureOpenAiChatOptions.builder().withDeploymentName(selectedModel).withMaxTokens(500).build()); + } + + @Bean + public String selectedModel() { + return Optional.ofNullable(System.getenv("AZURE_OPENAI_MODEL")).orElse("gpt-4-0125-preview"); } } diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/MockWeatherService.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/MockWeatherService.java index 898a1c61b..92747ed30 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/MockWeatherService.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/MockWeatherService.java @@ -15,14 +15,14 @@ */ package org.springframework.ai.azure.openai.function; -import java.util.function.Function; - import com.fasterxml.jackson.annotation.JsonClassDescription; import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; import com.fasterxml.jackson.annotation.JsonProperty; import com.fasterxml.jackson.annotation.JsonPropertyDescription; +import java.util.function.Function; + /** * @author Christian Tzolov */ @@ -87,4 +87,4 @@ public class MockWeatherService implements Function roleMap = new ConcurrentHashMap<>(); - return completionChunks.map(chunk -> toChatCompletion(chunk)).map(chatCompletion -> { + return completionChunks.map(chunk -> toChatCompletion(chunk)) + .switchMap( + cc -> handleFunctionCallOrReturnStream(request, Flux.just(ResponseEntity.of(Optional.of(cc))))) + .map(ResponseEntity::getBody) + .map(chatCompletion -> { + @SuppressWarnings("null") + String id = chatCompletion.id(); - chatCompletion = handleFunctionCallOrReturn(request, ResponseEntity.of(Optional.of(chatCompletion))) - .getBody(); - - @SuppressWarnings("null") - String id = chatCompletion.id(); - - List generations = chatCompletion.choices().stream().map(choice -> { - if (choice.message().role() != null) { - roleMap.putIfAbsent(id, choice.message().role().name()); - } - String finish = (choice.finishReason() != null ? choice.finishReason().name() : ""); - var generation = new Generation(choice.message().content(), - Map.of("id", id, "role", roleMap.get(id), "finishReason", finish)); - if (choice.finishReason() != null) { - generation = generation - .withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), null)); - } - return generation; - }).toList(); - return new ChatResponse(generations); - }); + List generations = chatCompletion.choices().stream().map(choice -> { + if (choice.message().role() != null) { + roleMap.putIfAbsent(id, choice.message().role().name()); + } + String finish = (choice.finishReason() != null ? choice.finishReason().name() : ""); + var generation = new Generation(choice.message().content(), + Map.of("id", id, "role", roleMap.get(id), "finishReason", finish)); + if (choice.finishReason() != null) { + generation = generation.withGenerationMetadata( + ChatGenerationMetadata.from(choice.finishReason().name(), null)); + } + return generation; + }).toList(); + return new ChatResponse(generations); + }); }); } @@ -271,7 +266,7 @@ public class MistralAiChatClient extends // Recursively call chatCompletionWithTools until the model doesn't call a // functions anymore. - ChatCompletionRequest newRequest = new ChatCompletionRequest(conversationHistory, false); + ChatCompletionRequest newRequest = new ChatCompletionRequest(conversationHistory, previousRequest.stream()); newRequest = ModelOptionsUtils.merge(newRequest, previousRequest, ChatCompletionRequest.class); return newRequest; @@ -299,6 +294,14 @@ public class MistralAiChatClient extends return this.mistralAiApi.chatCompletionEntity(request); } + @Override + protected Flux> doChatCompletionStream(ChatCompletionRequest request) { + return this.mistralAiApi.chatCompletionStream(request) + .map(this::toChatCompletion) + .map(Optional::ofNullable) + .map(ResponseEntity::of); + } + @Override protected boolean isToolFunctionCall(ResponseEntity chatCompletion) { diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java index 8e86ecdd1..6ec6904db 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java @@ -15,20 +15,8 @@ */ 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 reactor.core.publisher.Flux; - import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -57,6 +45,17 @@ 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 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; /** * {@link ChatClient} and {@link StreamingChatClient} implementation for {@literal OpenAI} @@ -68,6 +67,7 @@ import org.springframework.util.MimeType; * @author John Blum * @author Josh Long * @author Jemin Huh + * @author Grogdunn * @see ChatClient * @see StreamingChatClient * @see OpenAiApi @@ -189,36 +189,37 @@ public class OpenAiChatClient extends // Convert the ChatCompletionChunk into a ChatCompletion to be able to reuse // the function call handling logic. - return completionChunks.map(chunk -> chunkToChatCompletion(chunk)).map(chatCompletion -> { - try { - chatCompletion = handleFunctionCallOrReturn(request, ResponseEntity.of(Optional.of(chatCompletion))) - .getBody(); + return completionChunks.map(chunk -> chunkToChatCompletion(chunk)) + .switchMap( + cc -> handleFunctionCallOrReturnStream(request, Flux.just(ResponseEntity.of(Optional.of(cc))))) + .map(ResponseEntity::getBody) + .map(chatCompletion -> { + try { + @SuppressWarnings("null") + String id = chatCompletion.id(); - @SuppressWarnings("null") - String id = chatCompletion.id(); + List generations = chatCompletion.choices().stream().map(choice -> { + if (choice.message().role() != null) { + roleMap.putIfAbsent(id, choice.message().role().name()); + } + String finish = (choice.finishReason() != null ? choice.finishReason().name() : ""); + var generation = new Generation(choice.message().content(), + Map.of("id", id, "role", roleMap.get(id), "finishReason", finish)); + if (choice.finishReason() != null) { + generation = generation.withGenerationMetadata( + ChatGenerationMetadata.from(choice.finishReason().name(), null)); + } + return generation; + }).toList(); - List generations = chatCompletion.choices().stream().map(choice -> { - if (choice.message().role() != null) { - roleMap.putIfAbsent(id, choice.message().role().name()); - } - String finish = (choice.finishReason() != null ? choice.finishReason().name() : ""); - var generation = new Generation(choice.message().content(), - Map.of("id", id, "role", roleMap.get(id), "finishReason", finish)); - if (choice.finishReason() != null) { - generation = generation.withGenerationMetadata( - ChatGenerationMetadata.from(choice.finishReason().name(), null)); - } - return generation; - }).toList(); + return new ChatResponse(generations); + } + catch (Exception e) { + logger.error("Error processing chat completion", e); + return new ChatResponse(List.of()); + } - return new ChatResponse(generations); - } - catch (Exception e) { - logger.error("Error processing chat completion", e); - return new ChatResponse(List.of()); - } - - }); + }); }); } @@ -347,7 +348,7 @@ public class OpenAiChatClient extends // Recursively call chatCompletionWithTools until the model doesn't call a // functions anymore. - ChatCompletionRequest newRequest = new ChatCompletionRequest(conversationHistory, false); + ChatCompletionRequest newRequest = new ChatCompletionRequest(conversationHistory, previousRequest.stream()); newRequest = ModelOptionsUtils.merge(newRequest, previousRequest, ChatCompletionRequest.class); return newRequest; @@ -368,6 +369,14 @@ public class OpenAiChatClient extends 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(); diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java index ad74073fa..dfb2a14bc 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java @@ -15,12 +15,6 @@ */ package org.springframework.ai.vertexai.gemini; -import java.util.ArrayList; -import java.util.HashSet; -import java.util.List; -import java.util.Set; -import java.util.stream.Collectors; - import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; import com.google.cloud.vertexai.VertexAI; @@ -38,8 +32,6 @@ 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 reactor.core.publisher.Flux; - import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -60,9 +52,17 @@ 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 * @since 0.8.1 */ public class VertexAiGeminiChatClient @@ -167,18 +167,19 @@ public class VertexAiGeminiChatClient ResponseStream responseStream = request.model .generateContentStream(request.contents); - return Flux.fromStream(responseStream.stream()).map(response -> { - response = handleFunctionCallOrReturn(request, response); - List generations = response.getCandidatesList() - .stream() - .map(candidate -> candidate.getContent().getPartsList()) - .flatMap(List::stream) - .map(Part::getText) - .map(t -> new Generation(t.toString())) - .toList(); + 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.toString())) + .toList(); - return new ChatResponse(generations, toChatResponseMetadata(response)); - }); + return new ChatResponse(generations, toChatResponseMetadata(response)); + }); } catch (Exception e) { throw new RuntimeException("Failed to generate content", e); @@ -450,6 +451,19 @@ public class VertexAiGeminiChatClient } } + @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()) diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractFunctionCallSupport.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractFunctionCallSupport.java index d1c498625..d5be8ef6c 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractFunctionCallSupport.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractFunctionCallSupport.java @@ -15,6 +15,10 @@ */ package org.springframework.ai.model.function; +import org.springframework.util.CollectionUtils; +import reactor.core.publisher.Flux; +import reactor.core.publisher.Mono; + import java.util.ArrayList; import java.util.HashSet; import java.util.List; @@ -22,10 +26,9 @@ import java.util.Map; import java.util.Set; import java.util.concurrent.ConcurrentHashMap; -import org.springframework.util.CollectionUtils; - /** * @author Christian Tzolov + * @author Grogdunn */ public abstract class AbstractFunctionCallSupport { @@ -147,6 +150,36 @@ public abstract class AbstractFunctionCallSupport { return this.callWithFunctionSupport(newRequest); } + protected Flux callWithFunctionSupportStream(Req request) { + final Flux response = this.doChatCompletionStream(request); + return this.handleFunctionCallOrReturnStream(request, response); + } + + protected Flux handleFunctionCallOrReturnStream(Req request, Flux response) { + + return response.switchMap(resp -> { + if (!this.isToolFunctionCall(resp)) { + return Mono.just(resp); + } + + // The chat completion tool call requires the complete conversation + // history. Including the initial user message. + List conversationHistory = new ArrayList<>(); + + conversationHistory.addAll(this.doGetUserMessages(request)); + + Msg responseMessage = this.doGetToolResponseMessage(resp); + + // Add the assistant response to the message conversation history. + conversationHistory.add(responseMessage); + + Req newRequest = this.doCreateToolResponseRequest(request, responseMessage, conversationHistory); + + return this.callWithFunctionSupportStream(newRequest); + }); + + } + abstract protected Req doCreateToolResponseRequest(Req previousRequest, Msg responseMessage, List conversationHistory); @@ -156,6 +189,8 @@ public abstract class AbstractFunctionCallSupport { abstract protected Resp doChatCompletion(Req request); + abstract protected Flux doChatCompletionStream(Req request); + abstract protected boolean isToolFunctionCall(Resp response); }