From d42cd60eea9c8fd24b093a0ec037e62767590d12 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Mon, 25 Mar 2024 16:10:28 +0100 Subject: [PATCH] Fix: AbstractBedrockApi Streaming reusing of Flux Sink Resolves: #474 --- .../ai/bedrock/api/AbstractBedrockApi.java | 37 ++++++++++++------- .../BedrockAnthropicChatClientIT.java | 30 ++++++++++++++- .../BedrockAnthropicCreateRequestTests.java | 2 +- .../BedrockAnthropic3ChatClientIT.java | 29 +++++++++++++++ .../cohere/BedrockCohereChatClientIT.java | 29 +++++++++++++++ .../llama2/BedrockLlama2ChatClientIT.java | 29 +++++++++++++++ .../titan/BedrockTitanChatClientIT.java | 29 +++++++++++++++ 7 files changed, 169 insertions(+), 16 deletions(-) diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/api/AbstractBedrockApi.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/api/AbstractBedrockApi.java index 91c813e16..0618b6ef5 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/api/AbstractBedrockApi.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/api/AbstractBedrockApi.java @@ -18,6 +18,7 @@ package org.springframework.ai.bedrock.api; import java.io.UncheckedIOException; import java.nio.charset.StandardCharsets; +import java.time.Duration; import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; @@ -25,10 +26,12 @@ import com.fasterxml.jackson.annotation.JsonProperty; import com.fasterxml.jackson.core.JsonProcessingException; import com.fasterxml.jackson.databind.DeserializationFeature; import com.fasterxml.jackson.databind.ObjectMapper; -import org.apache.commons.logging.Log; -import org.apache.commons.logging.LogFactory; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; import reactor.core.publisher.Flux; import reactor.core.publisher.Sinks; +import reactor.core.publisher.Sinks.EmitFailureHandler; +import reactor.core.publisher.Sinks.EmitResult; import software.amazon.awssdk.auth.credentials.AwsCredentialsProvider; import software.amazon.awssdk.auth.credentials.ProfileCredentialsProvider; import software.amazon.awssdk.core.SdkBytes; @@ -60,7 +63,7 @@ import software.amazon.awssdk.services.bedrockruntime.model.ResponseStream; */ public abstract class AbstractBedrockApi { - private static final Log logger = LogFactory.getLog(AbstractBedrockApi.class); + private static final Logger logger = LoggerFactory.getLogger(AbstractBedrockApi.class); private final String modelId; private final ObjectMapper objectMapper; @@ -68,7 +71,6 @@ public abstract class AbstractBedrockApi { private final String region; private final BedrockRuntimeClient client; private final BedrockRuntimeAsyncClient clientStreaming; - private final Sinks.Many eventSink; /** * Create a new AbstractBedrockApi instance using default credentials provider and object mapper. @@ -96,8 +98,6 @@ public abstract class AbstractBedrockApi { this.credentialsProvider = credentialsProvider; this.region = region; - this.eventSink = Sinks.many().unicast().onBackpressureError(); - this.client = BedrockRuntimeClient.builder() .region(Region.of(this.region)) .credentialsProvider(this.credentialsProvider) @@ -220,13 +220,16 @@ public abstract class AbstractBedrockApi { */ protected Flux internalInvocationStream(I request, Class clazz) { + // final Sinks.Many eventSink = Sinks.many().unicast().onBackpressureError(); + final Sinks.Many eventSink = Sinks.many().multicast().onBackpressureBuffer(); + SdkBytes body; try { body = SdkBytes.fromUtf8String(this.objectMapper.writeValueAsString(request)); } catch (JsonProcessingException e) { - this.eventSink.tryEmitError(e); - return this.eventSink.asFlux(); + eventSink.tryEmitError(e); + return eventSink.asFlux(); } InvokeModelWithResponseStreamRequest invokeRequest = InvokeModelWithResponseStreamRequest.builder() @@ -240,16 +243,16 @@ public abstract class AbstractBedrockApi { try { logger.debug("Received chunk: " + chunk.bytes().asString(StandardCharsets.UTF_8)); SO response = this.objectMapper.readValue(chunk.bytes().asByteArray(), clazz); - this.eventSink.tryEmitNext(response); + eventSink.tryEmitNext(response); } catch (Exception e) { logger.error("Failed to unmarshall", e); - this.eventSink.tryEmitError(e); + eventSink.tryEmitError(e); } }) .onDefault((event) -> { logger.error("Unknown or unhandled event: " + event.toString()); - this.eventSink.tryEmitError(new Throwable("Unknown or unhandled event: " + event.toString())); + eventSink.tryEmitError(new Throwable("Unknown or unhandled event: " + event.toString())); }) .build(); @@ -257,12 +260,18 @@ public abstract class AbstractBedrockApi { .builder() .onComplete( () -> { - this.eventSink.tryEmitComplete(); + EmitResult emitResult = eventSink.tryEmitComplete(); + while(!emitResult.isSuccess()){ + System.out.println("Emitting complete:" + emitResult); + emitResult = eventSink.tryEmitComplete(); + }; + eventSink.emitComplete(EmitFailureHandler.busyLooping(Duration.ofSeconds(3))); + // EmitResult emitResult = eventSink.tryEmitComplete(); logger.debug("\nCompleted streaming response."); }) .onError((error) -> { logger.error("\n\nError streaming response: " + error.getMessage()); - this.eventSink.tryEmitError(error); + eventSink.tryEmitError(error); }) .onEventStream((stream) -> { stream.subscribe( @@ -274,7 +283,7 @@ public abstract class AbstractBedrockApi { this.clientStreaming.invokeModelWithResponseStream(invokeRequest, responseHandler); - return this.eventSink.asFlux(); + return eventSink.asFlux(); } } // @formatter:on \ No newline at end of file diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClientIT.java index 23d853aea..8935b47c0 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClientIT.java @@ -25,6 +25,7 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import reactor.core.publisher.Flux; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.AssistantMessage; @@ -64,6 +65,33 @@ class BedrockAnthropicChatClientIT { @Value("classpath:/prompts/system-message.st") private Resource systemResource; + @Test + void multipleStreamAttempts() { + + Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + + String joke1 = joke1Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + String joke2 = joke2Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + + assertThat(joke1).isNotBlank(); + assertThat(joke2).isNotBlank(); + } + @Test void roleTest() { UserMessage userMessage = new UserMessage( @@ -176,7 +204,7 @@ class BedrockAnthropicChatClientIT { @Bean public AnthropicChatBedrockApi anthropicApi() { return new AnthropicChatBedrockApi(AnthropicChatBedrockApi.AnthropicChatModel.CLAUDE_V2.id(), - EnvironmentVariableCredentialsProvider.create(), Region.EU_CENTRAL_1.id(), new ObjectMapper()); + EnvironmentVariableCredentialsProvider.create(), Region.US_EAST_1.id(), new ObjectMapper()); } @Bean diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicCreateRequestTests.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicCreateRequestTests.java index b37634792..aea62fc48 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicCreateRequestTests.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicCreateRequestTests.java @@ -32,7 +32,7 @@ import static org.assertj.core.api.Assertions.assertThat; public class BedrockAnthropicCreateRequestTests { private AnthropicChatBedrockApi anthropicChatApi = new AnthropicChatBedrockApi(AnthropicChatModel.CLAUDE_V2.id(), - Region.EU_CENTRAL_1.id()); + Region.US_EAST_1.id()); @Test public void createRequestWithChatOptions() { diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClientIT.java index 8048d5288..ed1af9804 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClientIT.java @@ -20,6 +20,8 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import reactor.core.publisher.Flux; + import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -67,6 +69,33 @@ class BedrockAnthropic3ChatClientIT { @Value("classpath:/prompts/system-message.st") private Resource systemResource; + @Test + void multipleStreamAttempts() { + + Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + + String joke1 = joke1Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + String joke2 = joke2Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + + assertThat(joke1).isNotBlank(); + assertThat(joke2).isNotBlank(); + } + @Test void roleTest() { UserMessage userMessage = new UserMessage( diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClientIT.java index 12512af50..f352ed456 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClientIT.java @@ -23,6 +23,8 @@ import java.util.stream.Collectors; import com.fasterxml.jackson.databind.ObjectMapper; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import reactor.core.publisher.Flux; + import org.springframework.ai.chat.messages.AssistantMessage; import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider; import software.amazon.awssdk.regions.Region; @@ -60,6 +62,33 @@ class BedrockCohereChatClientIT { @Value("classpath:/prompts/system-message.st") private Resource systemResource; + @Test + void multipleStreamAttempts() { + + Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + + String joke1 = joke1Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + String joke2 = joke2Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + + assertThat(joke1).isNotBlank(); + assertThat(joke2).isNotBlank(); + } + @Test void roleTest() { String request = "Tell me about 3 famous pirates from the Golden Age of Piracy and why they did."; diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama2/BedrockLlama2ChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama2/BedrockLlama2ChatClientIT.java index c73a3e926..09864c76a 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama2/BedrockLlama2ChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama2/BedrockLlama2ChatClientIT.java @@ -24,6 +24,8 @@ import com.fasterxml.jackson.databind.ObjectMapper; import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import reactor.core.publisher.Flux; + import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.AssistantMessage; import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider; @@ -61,6 +63,33 @@ class BedrockLlama2ChatClientIT { @Value("classpath:/prompts/system-message.st") private Resource systemResource; + @Test + void multipleStreamAttempts() { + + Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + + String joke1 = joke1Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + String joke2 = joke2Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + + assertThat(joke1).isNotBlank(); + assertThat(joke2).isNotBlank(); + } + @Test void roleTest() { UserMessage userMessage = new UserMessage( diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClientIT.java index 51fae542c..d8a04dd7f 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClientIT.java @@ -24,6 +24,8 @@ import com.fasterxml.jackson.databind.ObjectMapper; import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import reactor.core.publisher.Flux; + import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.AssistantMessage; import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider; @@ -61,6 +63,33 @@ class BedrockTitanChatClientIT { @Value("classpath:/prompts/system-message.st") private Resource systemResource; + @Test + void multipleStreamAttempts() { + + Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + + String joke1 = joke1Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + String joke2 = joke2Stream.collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + + assertThat(joke1).isNotBlank(); + assertThat(joke2).isNotBlank(); + } + @Test void roleTest() { String request = "Tell me about 3 famous pirates from the Golden Age of Piracy and why they did.";