From 178a607cf6fc65e302b7420fe50cb8dff7e2df2d Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Mon, 27 May 2024 00:51:36 +0200 Subject: [PATCH] Fix an issue related to the Azure OpenAI streamig response This issue in the Flux streaming logic was causing chunks to be aggreagated in a single response. Re-enable the AzureOpenAiAutoConfigurationIT --- .../ai/azure/openai/AzureOpenAiChatModel.java | 4 ++-- .../ai/openai/chat/OpenAiChatModelIT.java | 23 +++++++++++++++++++ .../azure/AzureOpenAiAutoConfigurationIT.java | 2 -- .../tool/FunctionCallWithFunctionBeanIT.java | 1 - 4 files changed, 25 insertions(+), 5 deletions(-) diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java index dd53e9503..9ac782f6c 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java @@ -183,8 +183,8 @@ public class AzureOpenAiChatModel isFunctionCall.set(false); return true; } - return false; - }, false) + return !isFunctionCall.get(); + }) .concatMapIterable(window -> { final var reduce = window.reduce(MergeUtils.emptyChatCompletions(), MergeUtils::mergeChatCompletions); return List.of(reduce); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java index 3d5cd08d3..b0adcc1e0 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java @@ -80,6 +80,29 @@ class OpenAiChatModelIT extends AbstractIT { // needs fine tuning... evaluateQuestionAndAnswer(request, response, false); } + @Test + void streamRoleTest() { + UserMessage userMessage = new UserMessage( + "Tell me about 3 famous pirates from the Golden Age of Piracy and what they did."); + SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource); + Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", "Bob", "voice", "pirate")); + Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); + Flux flux = streamingChatModel.stream(prompt); + + List responses = flux.collectList().block(); + assertThat(responses.size()).isGreaterThan(1); + + String stitchedResponseContent = responses.stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + + assertThat(stitchedResponseContent).contains("Blackbeard"); + + } + @Test void listOutputConverter() { DefaultConversionService conversionService = new DefaultConversionService(); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/AzureOpenAiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/AzureOpenAiAutoConfigurationIT.java index c188b73a2..0f514491b 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/AzureOpenAiAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/AzureOpenAiAutoConfigurationIT.java @@ -19,7 +19,6 @@ import java.util.List; import java.util.Map; import java.util.stream.Collectors; -import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.azure.openai.AzureOpenAiChatModel; @@ -44,7 +43,6 @@ import static org.assertj.core.api.Assertions.assertThat; * @author Christian Tzolov * @since 0.8.0 */ -@Disabled("streaming response on mark p machine is not returning a list of size > 1") @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") public class AzureOpenAiAutoConfigurationIT { diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java index 3cb6b1567..799a37623 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java @@ -35,7 +35,6 @@ import org.springframework.boot.test.context.runner.ApplicationContextRunner; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; import org.springframework.context.annotation.Description; -import org.springframework.util.StringUtils; import static org.assertj.core.api.Assertions.assertThat; import static org.springframework.ai.autoconfigure.azure.tool.DeploymentNameUtil.getDeploymentName;