fix: Improve error resilience of the Bedrock stream handling

The changes include:

- Refactor the emitting of next, error, and complete events in the Bedrock stream handling to use a default EmitFailureHandler that retries for 10 seconds before failing.
  This helps improve the error resilience of the stream processing.
- Disable a couple of integration tests related to the COHERE_COMMAND_V14 model, as that model version is no longer supported.
- Adjust some configuration options in the Jurassic2 chat model integration test.

Also resolves #1679
This commit is contained in:
Christian Tzolov
2024-11-08 13:53:33 +01:00
committed by Ilayaperumal Gopinathan
parent 7895875ca2
commit 9fa89a2fd6
6 changed files with 25 additions and 28 deletions

View File

@@ -36,7 +36,6 @@ import reactor.core.publisher.Flux;
import reactor.core.publisher.Mono;
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.core.SdkBytes;
import software.amazon.awssdk.core.document.Document;
@@ -524,6 +523,9 @@ public class BedrockProxyChatModel extends AbstractToolCallSupport implements Ch
});
}
public static final EmitFailureHandler DEFAULT_EMIT_FAILURE_HANDLER = EmitFailureHandler
.busyLooping(Duration.ofSeconds(10));
/**
* Invoke the model and return the response stream.
*
@@ -541,26 +543,19 @@ public class BedrockProxyChatModel extends AbstractToolCallSupport implements Ch
ConverseStreamResponseHandler.Visitor visitor = ConverseStreamResponseHandler.Visitor.builder()
.onDefault(output -> {
logger.debug("Received converse stream output:{}", output);
eventSink.tryEmitNext(output);
eventSink.emitNext(output, DEFAULT_EMIT_FAILURE_HANDLER);
})
.build();
ConverseStreamResponseHandler responseHandler = ConverseStreamResponseHandler.builder()
.onEventStream(stream -> stream.subscribe(e -> e.accept(visitor)))
.onComplete(() -> {
EmitResult emitResult = eventSink.tryEmitComplete();
while (!emitResult.isSuccess()) {
logger.info("Emitting complete:{}", emitResult);
emitResult = eventSink.tryEmitComplete();
}
eventSink.emitComplete(EmitFailureHandler.busyLooping(Duration.ofSeconds(3)));
eventSink.emitComplete(DEFAULT_EMIT_FAILURE_HANDLER);
logger.info("Completed streaming response.");
})
.onError(error -> {
logger.error("Error handling Bedrock converse stream response", error);
eventSink.tryEmitError(error);
eventSink.emitError(error, DEFAULT_EMIT_FAILURE_HANDLER);
})
.build();

View File

@@ -32,7 +32,6 @@ 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;
@@ -70,6 +69,10 @@ public abstract class AbstractBedrockApi<I, O, SO> {
private static final Logger logger = LoggerFactory.getLogger(AbstractBedrockApi.class);
public static final EmitFailureHandler DEFAULT_EMIT_FAILURE_HANDLER = EmitFailureHandler
.busyLooping(Duration.ofSeconds(10));
private final String modelId;
private final ObjectMapper objectMapper;
private final Region region;
@@ -264,7 +267,7 @@ public abstract class AbstractBedrockApi<I, O, SO> {
body = SdkBytes.fromUtf8String(this.objectMapper.writeValueAsString(request));
}
catch (JsonProcessingException e) {
eventSink.tryEmitError(e);
eventSink.emitError(e, DEFAULT_EMIT_FAILURE_HANDLER);
return eventSink.asFlux();
}
@@ -279,16 +282,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);
eventSink.tryEmitNext(response);
eventSink.emitNext(response, DEFAULT_EMIT_FAILURE_HANDLER);
}
catch (Exception e) {
logger.error("Failed to unmarshall", e);
eventSink.tryEmitError(e);
eventSink.emitError(e, DEFAULT_EMIT_FAILURE_HANDLER);
}
})
.onDefault(event -> {
logger.error("Unknown or unhandled event: " + event.toString());
eventSink.tryEmitError(new Throwable("Unknown or unhandled event: " + event.toString()));
eventSink.emitError(new Throwable("Unknown or unhandled event: " + event.toString()),DEFAULT_EMIT_FAILURE_HANDLER);
})
.build();
@@ -296,18 +299,12 @@ public abstract class AbstractBedrockApi<I, O, SO> {
.builder()
.onComplete(
() -> {
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.");
eventSink.emitComplete(DEFAULT_EMIT_FAILURE_HANDLER);
logger.info("Completed streaming response.");
})
.onError(error -> {
logger.error("\n\nError streaming response: " + error.getMessage());
eventSink.tryEmitError(error);
eventSink.emitError(error, DEFAULT_EMIT_FAILURE_HANDLER);
})
.onEventStream(stream -> stream.subscribe(
(ResponseStream e) -> e.accept(visitor)))

View File

@@ -23,6 +23,7 @@ import java.util.Map;
import java.util.stream.Collectors;
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;
@@ -55,6 +56,7 @@ import static org.assertj.core.api.Assertions.assertThat;
@SpringBootTest
@EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*")
@EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*")
@Disabled("COHERE_COMMAND_V14 is not supported anymore")
class BedrockCohereChatModelIT {
@Autowired

View File

@@ -19,6 +19,7 @@ package org.springframework.ai.bedrock.cohere.api;
import java.time.Duration;
import java.util.List;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import reactor.core.publisher.Flux;
@@ -71,6 +72,7 @@ public class CohereChatBedrockApiIT {
}
@Test
@Disabled("Due to model version has reached the end of its life")
public void chatCompletion() {
var request = CohereChatRequest
@@ -95,6 +97,7 @@ public class CohereChatBedrockApiIT {
assertThat(response.generations().get(0).text()).isNotEmpty();
}
@Disabled("Due to model version has reached the end of its life")
@Test
public void chatCompletionStream() {

View File

@@ -158,8 +158,8 @@ class BedrockAi21Jurassic2ChatModelIT {
return new BedrockAi21Jurassic2ChatModel(jurassic2ChatBedrockApi,
BedrockAi21Jurassic2ChatOptions.builder()
.withTemperature(0.5)
.withMaxTokens(100)
.withTopP(0.9)
.withMaxTokens(500)
// .withTopP(0.9)
.build());
}

View File

@@ -51,7 +51,7 @@ import org.springframework.context.annotation.Import;
*/
@AutoConfiguration
@EnableConfigurationProperties({ BedrockConverseProxyChatProperties.class, BedrockAwsConnectionConfiguration.class })
@ConditionalOnClass({ BedrockRuntimeClient.class, BedrockRuntimeAsyncClient.class })
@ConditionalOnClass({ BedrockProxyChatModel.class, BedrockRuntimeClient.class, BedrockRuntimeAsyncClient.class })
@ConditionalOnProperty(prefix = BedrockConverseProxyChatProperties.CONFIG_PREFIX, name = "enabled",
havingValue = "true", matchIfMissing = true)
@Import(BedrockAwsConnectionConfiguration.class)