use aws_sdk_bedrockruntime::error::SdkError;
use aws_sdk_bedrockruntime::primitives::event_stream::EventReceiver;
use aws_sdk_bedrockruntime::types::error::ConverseStreamOutputError;
use aws_sdk_bedrockruntime::types::{
ContentBlockDelta, ContentBlockStart, ConverseStreamOutput, ReasoningContentBlockDelta,
StopReason as BedrockStopReason, TokenUsage as BedrockTokenUsage,
};
use aws_smithy_types::event_stream::RawMessage;
use futures::{Stream, stream};
use std::future::ready;
use tracing::{error, warn};
use crate::provider_connection::DEFAULT_STREAM_IDLE_TIMEOUT;
use crate::providers::response_stream::{OpenedStream, StreamAssembler, response_stream};
use crate::{LlmError, LlmResponse, LlmResponseStream, ProviderError, StopReason, TokenUsage, Tokens};
pub fn process_bedrock_stream(
events: impl Stream<Item = crate::Result<ConverseStreamOutput>> + Send + 'static,
) -> LlmResponseStream {
response_stream(
ready(Ok(OpenedStream::new(events))),
|event, turn| Ok(decode_converse_event(event, turn)),
DEFAULT_STREAM_IDLE_TIMEOUT,
)
}
pub(crate) fn converse_events(
receiver: EventReceiver<ConverseStreamOutput, ConverseStreamOutputError>,
) -> impl Stream<Item = crate::Result<ConverseStreamOutput>> + Send {
stream::unfold(receiver, |mut receiver| async move {
let event = receiver.recv().await.map_err(|e| {
error!("Bedrock stream recv error: {e}");
LlmError::from(e)
});
event.transpose().map(|event| (event, receiver))
})
}
impl From<&BedrockTokenUsage> for TokenUsage {
fn from(usage: &BedrockTokenUsage) -> Self {
let cache_read = usage.cache_read_input_tokens().and_then(|v| u32::try_from(v).ok()).map(Tokens::from);
let cache_creation = usage.cache_write_input_tokens().and_then(|v| u32::try_from(v).ok()).map(Tokens::from);
let cached = cache_read.unwrap_or_default() + cache_creation.unwrap_or_default();
TokenUsage {
input_tokens: Tokens::from(u32::try_from(usage.input_tokens).unwrap_or(0)) + cached,
output_tokens: u32::try_from(usage.output_tokens).unwrap_or(0).into(),
cache_read_tokens: cache_read,
cache_creation_tokens: cache_creation,
..TokenUsage::default()
}
}
}
impl From<SdkError<ConverseStreamOutputError, RawMessage>> for LlmError {
fn from(e: SdkError<ConverseStreamOutputError, RawMessage>) -> Self {
let message = format!("Bedrock stream error: {e}");
let provider = match e {
SdkError::ServiceError(svc) => {
let inner = svc.err();
if inner.is_throttling_exception() {
ProviderError::rate_limit(message)
} else if inner.is_service_unavailable_exception()
|| inner.is_internal_server_exception()
|| inner.is_model_stream_error_exception()
{
ProviderError::stream_interrupted(message)
} else {
ProviderError::api(message)
}
}
_ => ProviderError::stream_interrupted(message),
};
Self::from(provider)
}
}
pub(super) fn decode_converse_event(event: ConverseStreamOutput, turn: &mut StreamAssembler<i32>) -> Vec<LlmResponse> {
let response = match event {
ConverseStreamOutput::ContentBlockStart(event) => match event.start {
Some(ContentBlockStart::ToolUse(tool)) => {
Some(turn.start_tool(event.content_block_index, tool.tool_use_id, tool.name))
}
_ => None,
},
ConverseStreamOutput::ContentBlockDelta(event) => match event.delta {
Some(ContentBlockDelta::Text(text)) if !text.is_empty() => Some(LlmResponse::Text { chunk: text }),
Some(ContentBlockDelta::ToolUse(delta)) => turn.append_tool_args(&event.content_block_index, delta.input),
Some(ContentBlockDelta::ReasoningContent(ReasoningContentBlockDelta::Text(text))) if !text.is_empty() => {
Some(LlmResponse::Reasoning { chunk: text })
}
_ => None,
},
ConverseStreamOutput::ContentBlockStop(event) => turn.complete_tool(&event.content_block_index),
ConverseStreamOutput::MessageStop(event) => {
turn.stop(map_bedrock_stop_reason(&event.stop_reason));
turn.allow_eof();
None
}
ConverseStreamOutput::Metadata(event) => {
event.usage.as_ref().map(|usage| LlmResponse::Usage { tokens: usage.into() })
}
ConverseStreamOutput::MessageStart(_) => None,
other => {
warn!("Unhandled Bedrock stream event: {other:?}");
None
}
};
response.into_iter().collect()
}
fn map_bedrock_stop_reason(reason: &BedrockStopReason) -> StopReason {
match reason {
BedrockStopReason::EndTurn | BedrockStopReason::StopSequence => StopReason::EndTurn,
BedrockStopReason::ToolUse => StopReason::ToolCalls,
BedrockStopReason::MaxTokens | BedrockStopReason::ModelContextWindowExceeded => StopReason::Length,
BedrockStopReason::ContentFiltered | BedrockStopReason::GuardrailIntervened => StopReason::ContentFilter,
other => StopReason::Unknown(format!("{other:?}")),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ProviderErrorKind;
use crate::testing::llm_response;
use aws_sdk_bedrockruntime::types::{
ContentBlockDeltaEvent, ContentBlockStartEvent, ContentBlockStopEvent, ConversationRole,
ConverseStreamMetadataEvent, MessageStartEvent, MessageStopEvent, ToolUseBlockDelta, ToolUseBlockStart,
};
use futures::StreamExt;
#[tokio::test]
async fn test_text_stream() {
let responses = collect_responses(
bedrock_stream().text(0, &["Hello", "", " world"]).message_stop(BedrockStopReason::EndTurn).build(),
)
.await;
assert_eq!(responses, llm_response().text(&["Hello", " world"]).build_with_stop_reason(StopReason::EndTurn));
}
#[tokio::test]
async fn test_reasoning_stream() {
let responses = collect_responses(
bedrock_stream()
.reasoning(0, &["thinking"])
.text(1, &["answer"])
.message_stop(BedrockStopReason::EndTurn)
.build(),
)
.await;
assert_eq!(
responses,
llm_response().reasoning(&["thinking"]).text(&["answer"]).build_with_stop_reason(StopReason::EndTurn)
);
}
#[tokio::test]
async fn test_tool_call_stream() {
let deltas = [r#"{"query":"#, r#""test"}"#];
let responses = collect_responses(
bedrock_stream()
.tool_call(0, "tool_123", "search", &deltas)
.message_stop(BedrockStopReason::ToolUse)
.build(),
)
.await;
assert_eq!(
responses,
llm_response().tool_call("tool_123", "search", &deltas).build_with_stop_reason(StopReason::ToolCalls)
);
}
#[tokio::test]
async fn stream_closed_mid_tool_call_does_not_complete_it() {
let responses =
process_events(bedrock_stream().tool_start(0, "tool_123", "search").tool_delta(0, r#"{"query":"#).build())
.await;
assert!(
matches!(
responses.as_slice(),
[
Ok(LlmResponse::Start),
Ok(LlmResponse::ToolRequestStart { .. }),
Ok(LlmResponse::ToolRequestArg { .. }),
Err(error)
] if error.provider().map(|provider| provider.kind) == Some(ProviderErrorKind::StreamInterrupted)
),
"{responses:?}"
);
}
#[tokio::test]
async fn test_metadata_after_message_stop_reports_cache_usage() {
let usage = BedrockTokenUsage::builder()
.input_tokens(100)
.output_tokens(50)
.total_tokens(150)
.cache_read_input_tokens(40)
.cache_write_input_tokens(20)
.build()
.unwrap();
let responses =
collect_responses(bedrock_stream().message_stop(BedrockStopReason::EndTurn).metadata(usage).build()).await;
let usage = responses.iter().find_map(|response| match response {
LlmResponse::Usage { tokens } => Some(*tokens),
_ => None,
});
assert_eq!(
usage,
Some(TokenUsage {
input_tokens: 160.into(),
output_tokens: 50.into(),
cache_read_tokens: Some(40.into()),
cache_creation_tokens: Some(20.into()),
..TokenUsage::default()
}),
"cached tokens count toward the prompt"
);
assert_eq!(responses.last(), Some(&LlmResponse::done_with_stop_reason(StopReason::EndTurn)));
}
#[tokio::test]
async fn test_metadata_without_cache_fields() {
let usage = BedrockTokenUsage::builder().input_tokens(10).output_tokens(5).total_tokens(15).build().unwrap();
let responses =
collect_responses(bedrock_stream().message_stop(BedrockStopReason::EndTurn).metadata(usage).build()).await;
assert_eq!(responses, llm_response().usage(10, 5).build_with_stop_reason(StopReason::EndTurn));
}
#[tokio::test]
async fn test_stop_reasons_map_to_llm_stop_reasons() {
for (bedrock_stop_reason, stop_reason) in [
(BedrockStopReason::EndTurn, StopReason::EndTurn),
(BedrockStopReason::StopSequence, StopReason::EndTurn),
(BedrockStopReason::ToolUse, StopReason::ToolCalls),
(BedrockStopReason::MaxTokens, StopReason::Length),
(BedrockStopReason::ModelContextWindowExceeded, StopReason::Length),
(BedrockStopReason::ContentFiltered, StopReason::ContentFilter),
(BedrockStopReason::GuardrailIntervened, StopReason::ContentFilter),
] {
let responses = collect_responses(bedrock_stream().message_stop(bedrock_stop_reason).build()).await;
assert_eq!(responses, llm_response().build_with_stop_reason(stop_reason));
}
}
async fn collect_responses(events: Vec<ConverseStreamOutput>) -> Vec<LlmResponse> {
process_events(events).await.into_iter().map(Result::unwrap).collect()
}
async fn process_events(events: Vec<ConverseStreamOutput>) -> Vec<crate::Result<LlmResponse>> {
process_bedrock_stream(stream::iter(events.into_iter().map(Ok))).collect().await
}
fn bedrock_stream() -> BedrockStreamBuilder {
BedrockStreamBuilder::default().push(ConverseStreamOutput::MessageStart(
MessageStartEvent::builder().role(ConversationRole::Assistant).build().unwrap(),
))
}
#[derive(Default)]
struct BedrockStreamBuilder {
events: Vec<ConverseStreamOutput>,
}
impl BedrockStreamBuilder {
fn text(self, index: i32, chunks: &[&str]) -> Self {
chunks
.iter()
.fold(self, |builder, chunk| builder.delta(index, ContentBlockDelta::Text((*chunk).to_string())))
.block_stop(index)
}
fn reasoning(self, index: i32, chunks: &[&str]) -> Self {
chunks
.iter()
.fold(self, |builder, chunk| {
builder.delta(
index,
ContentBlockDelta::ReasoningContent(ReasoningContentBlockDelta::Text((*chunk).to_string())),
)
})
.block_stop(index)
}
fn tool_call(self, index: i32, id: &str, name: &str, argument_deltas: &[&str]) -> Self {
argument_deltas
.iter()
.fold(self.tool_start(index, id, name), |builder, delta| builder.tool_delta(index, delta))
.block_stop(index)
}
fn tool_start(self, index: i32, id: &str, name: &str) -> Self {
let tool = ToolUseBlockStart::builder().tool_use_id(id).name(name).build().unwrap();
self.push(ConverseStreamOutput::ContentBlockStart(
ContentBlockStartEvent::builder()
.content_block_index(index)
.start(ContentBlockStart::ToolUse(tool))
.build()
.unwrap(),
))
}
fn tool_delta(self, index: i32, input: &str) -> Self {
self.delta(index, ContentBlockDelta::ToolUse(ToolUseBlockDelta::builder().input(input).build().unwrap()))
}
fn message_stop(self, stop_reason: BedrockStopReason) -> Self {
self.push(ConverseStreamOutput::MessageStop(
MessageStopEvent::builder().stop_reason(stop_reason).build().unwrap(),
))
}
fn metadata(self, usage: BedrockTokenUsage) -> Self {
self.push(ConverseStreamOutput::Metadata(ConverseStreamMetadataEvent::builder().usage(usage).build()))
}
fn delta(self, index: i32, delta: ContentBlockDelta) -> Self {
self.push(ConverseStreamOutput::ContentBlockDelta(
ContentBlockDeltaEvent::builder().content_block_index(index).delta(delta).build().unwrap(),
))
}
fn block_stop(self, index: i32) -> Self {
self.push(ConverseStreamOutput::ContentBlockStop(
ContentBlockStopEvent::builder().content_block_index(index).build().unwrap(),
))
}
fn push(mut self, event: ConverseStreamOutput) -> Self {
self.events.push(event);
self
}
fn build(self) -> Vec<ConverseStreamOutput> {
self.events
}
}
}