Fix: AbstractBedrockApi Streaming reusing of Flux Sink

Resolves: #474
This commit is contained in:
Christian Tzolov
2024-03-25 16:10:28 +01:00
parent 34c4703ea0
commit d42cd60eea
7 changed files with 169 additions and 16 deletions

View File

@@ -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<I, O, SO> {
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<I, O, SO> {
private final String region;
private final BedrockRuntimeClient client;
private final BedrockRuntimeAsyncClient clientStreaming;
private final Sinks.Many<SO> eventSink;
/**
* Create a new AbstractBedrockApi instance using default credentials provider and object mapper.
@@ -96,8 +98,6 @@ public abstract class AbstractBedrockApi<I, O, SO> {
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<I, O, SO> {
*/
protected Flux<SO> internalInvocationStream(I request, Class<SO> clazz) {
// final Sinks.Many<SO> eventSink = Sinks.many().unicast().onBackpressureError();
final Sinks.Many<SO> 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<I, O, SO> {
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<I, O, SO> {
.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<I, O, SO> {
this.clientStreaming.invokeModelWithResponseStream(invokeRequest, responseHandler);
return this.eventSink.asFlux();
return eventSink.asFlux();
}
}
// @formatter:on

View File

@@ -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<ChatResponse> joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?")));
Flux<ChatResponse> 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

View File

@@ -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() {

View File

@@ -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<ChatResponse> joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?")));
Flux<ChatResponse> 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(

View File

@@ -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<ChatResponse> joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?")));
Flux<ChatResponse> 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.";

View File

@@ -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<ChatResponse> joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?")));
Flux<ChatResponse> 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(

View File

@@ -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<ChatResponse> joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?")));
Flux<ChatResponse> 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.";