@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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.";
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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.";
|
||||
|
||||
Reference in New Issue
Block a user