Skip to main content

llm/providers/bedrock/
streaming.rs

1use aws_sdk_bedrockruntime::error::SdkError;
2use aws_sdk_bedrockruntime::primitives::event_stream::EventReceiver;
3use aws_sdk_bedrockruntime::types::error::ConverseStreamOutputError;
4use aws_sdk_bedrockruntime::types::{
5    ContentBlockDelta, ContentBlockStart, ConverseStreamOutput, ReasoningContentBlockDelta,
6    StopReason as BedrockStopReason, TokenUsage as BedrockTokenUsage,
7};
8use aws_smithy_types::event_stream::RawMessage;
9use futures::{Stream, stream};
10use std::future::ready;
11use tracing::{error, warn};
12
13use crate::provider_connection::DEFAULT_STREAM_IDLE_TIMEOUT;
14
15use crate::providers::response_stream::{OpenedStream, StreamAssembler, response_stream};
16use crate::{LlmError, LlmResponse, LlmResponseStream, ProviderError, StopReason, TokenUsage, Tokens};
17
18pub fn process_bedrock_stream(
19    events: impl Stream<Item = crate::Result<ConverseStreamOutput>> + Send + 'static,
20) -> LlmResponseStream {
21    response_stream(
22        ready(Ok(OpenedStream::new(events))),
23        |event, turn| Ok(decode_converse_event(event, turn)),
24        DEFAULT_STREAM_IDLE_TIMEOUT,
25    )
26}
27
28pub(crate) fn converse_events(
29    receiver: EventReceiver<ConverseStreamOutput, ConverseStreamOutputError>,
30) -> impl Stream<Item = crate::Result<ConverseStreamOutput>> + Send {
31    stream::unfold(receiver, |mut receiver| async move {
32        let event = receiver.recv().await.map_err(|e| {
33            error!("Bedrock stream recv error: {e}");
34            LlmError::from(e)
35        });
36        event.transpose().map(|event| (event, receiver))
37    })
38}
39
40impl From<&BedrockTokenUsage> for TokenUsage {
41    fn from(usage: &BedrockTokenUsage) -> Self {
42        let cache_read = usage.cache_read_input_tokens().and_then(|v| u32::try_from(v).ok()).map(Tokens::from);
43        let cache_creation = usage.cache_write_input_tokens().and_then(|v| u32::try_from(v).ok()).map(Tokens::from);
44        // Bedrock's input_tokens excludes cached tokens; TokenUsage counts the whole prompt.
45        let cached = cache_read.unwrap_or_default() + cache_creation.unwrap_or_default();
46        TokenUsage {
47            input_tokens: Tokens::from(u32::try_from(usage.input_tokens).unwrap_or(0)) + cached,
48            output_tokens: u32::try_from(usage.output_tokens).unwrap_or(0).into(),
49            cache_read_tokens: cache_read,
50            cache_creation_tokens: cache_creation,
51            ..TokenUsage::default()
52        }
53    }
54}
55
56impl From<SdkError<ConverseStreamOutputError, RawMessage>> for LlmError {
57    fn from(e: SdkError<ConverseStreamOutputError, RawMessage>) -> Self {
58        let message = format!("Bedrock stream error: {e}");
59        let provider = match e {
60            SdkError::ServiceError(svc) => {
61                let inner = svc.err();
62                if inner.is_throttling_exception() {
63                    ProviderError::rate_limit(message)
64                } else if inner.is_service_unavailable_exception()
65                    || inner.is_internal_server_exception()
66                    || inner.is_model_stream_error_exception()
67                {
68                    ProviderError::stream_interrupted(message)
69                } else {
70                    ProviderError::api(message)
71                }
72            }
73            _ => ProviderError::stream_interrupted(message),
74        };
75        Self::from(provider)
76    }
77}
78
79pub(super) fn decode_converse_event(event: ConverseStreamOutput, turn: &mut StreamAssembler<i32>) -> Vec<LlmResponse> {
80    let response = match event {
81        ConverseStreamOutput::ContentBlockStart(event) => match event.start {
82            Some(ContentBlockStart::ToolUse(tool)) => {
83                Some(turn.start_tool(event.content_block_index, tool.tool_use_id, tool.name))
84            }
85            _ => None,
86        },
87        ConverseStreamOutput::ContentBlockDelta(event) => match event.delta {
88            Some(ContentBlockDelta::Text(text)) if !text.is_empty() => Some(LlmResponse::Text { chunk: text }),
89            Some(ContentBlockDelta::ToolUse(delta)) => turn.append_tool_args(&event.content_block_index, delta.input),
90            Some(ContentBlockDelta::ReasoningContent(ReasoningContentBlockDelta::Text(text))) if !text.is_empty() => {
91                Some(LlmResponse::Reasoning { chunk: text })
92            }
93            _ => None,
94        },
95        ConverseStreamOutput::ContentBlockStop(event) => turn.complete_tool(&event.content_block_index),
96        ConverseStreamOutput::MessageStop(event) => {
97            turn.stop(map_bedrock_stop_reason(&event.stop_reason));
98            turn.allow_eof();
99            None
100        }
101        ConverseStreamOutput::Metadata(event) => {
102            event.usage.as_ref().map(|usage| LlmResponse::Usage { tokens: usage.into() })
103        }
104        ConverseStreamOutput::MessageStart(_) => None,
105        other => {
106            warn!("Unhandled Bedrock stream event: {other:?}");
107            None
108        }
109    };
110
111    response.into_iter().collect()
112}
113
114fn map_bedrock_stop_reason(reason: &BedrockStopReason) -> StopReason {
115    match reason {
116        BedrockStopReason::EndTurn | BedrockStopReason::StopSequence => StopReason::EndTurn,
117        BedrockStopReason::ToolUse => StopReason::ToolCalls,
118        BedrockStopReason::MaxTokens | BedrockStopReason::ModelContextWindowExceeded => StopReason::Length,
119        BedrockStopReason::ContentFiltered | BedrockStopReason::GuardrailIntervened => StopReason::ContentFilter,
120        other => StopReason::Unknown(format!("{other:?}")),
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    use super::*;
127    use crate::ProviderErrorKind;
128    use crate::testing::llm_response;
129    use aws_sdk_bedrockruntime::types::{
130        ContentBlockDeltaEvent, ContentBlockStartEvent, ContentBlockStopEvent, ConversationRole,
131        ConverseStreamMetadataEvent, MessageStartEvent, MessageStopEvent, ToolUseBlockDelta, ToolUseBlockStart,
132    };
133    use futures::StreamExt;
134
135    #[tokio::test]
136    async fn test_text_stream() {
137        let responses = collect_responses(
138            bedrock_stream().text(0, &["Hello", "", " world"]).message_stop(BedrockStopReason::EndTurn).build(),
139        )
140        .await;
141
142        assert_eq!(responses, llm_response().text(&["Hello", " world"]).build_with_stop_reason(StopReason::EndTurn));
143    }
144
145    #[tokio::test]
146    async fn test_reasoning_stream() {
147        let responses = collect_responses(
148            bedrock_stream()
149                .reasoning(0, &["thinking"])
150                .text(1, &["answer"])
151                .message_stop(BedrockStopReason::EndTurn)
152                .build(),
153        )
154        .await;
155
156        assert_eq!(
157            responses,
158            llm_response().reasoning(&["thinking"]).text(&["answer"]).build_with_stop_reason(StopReason::EndTurn)
159        );
160    }
161
162    #[tokio::test]
163    async fn test_tool_call_stream() {
164        let deltas = [r#"{"query":"#, r#""test"}"#];
165
166        let responses = collect_responses(
167            bedrock_stream()
168                .tool_call(0, "tool_123", "search", &deltas)
169                .message_stop(BedrockStopReason::ToolUse)
170                .build(),
171        )
172        .await;
173
174        assert_eq!(
175            responses,
176            llm_response().tool_call("tool_123", "search", &deltas).build_with_stop_reason(StopReason::ToolCalls)
177        );
178    }
179
180    #[tokio::test]
181    async fn stream_closed_mid_tool_call_does_not_complete_it() {
182        let responses =
183            process_events(bedrock_stream().tool_start(0, "tool_123", "search").tool_delta(0, r#"{"query":"#).build())
184                .await;
185
186        assert!(
187            matches!(
188                responses.as_slice(),
189                [
190                    Ok(LlmResponse::Start),
191                    Ok(LlmResponse::ToolRequestStart { .. }),
192                    Ok(LlmResponse::ToolRequestArg { .. }),
193                    Err(error)
194                ] if error.provider().map(|provider| provider.kind) == Some(ProviderErrorKind::StreamInterrupted)
195            ),
196            "{responses:?}"
197        );
198    }
199
200    #[tokio::test]
201    async fn test_metadata_after_message_stop_reports_cache_usage() {
202        let usage = BedrockTokenUsage::builder()
203            .input_tokens(100)
204            .output_tokens(50)
205            .total_tokens(150)
206            .cache_read_input_tokens(40)
207            .cache_write_input_tokens(20)
208            .build()
209            .unwrap();
210
211        let responses =
212            collect_responses(bedrock_stream().message_stop(BedrockStopReason::EndTurn).metadata(usage).build()).await;
213
214        let usage = responses.iter().find_map(|response| match response {
215            LlmResponse::Usage { tokens } => Some(*tokens),
216            _ => None,
217        });
218        assert_eq!(
219            usage,
220            Some(TokenUsage {
221                input_tokens: 160.into(),
222                output_tokens: 50.into(),
223                cache_read_tokens: Some(40.into()),
224                cache_creation_tokens: Some(20.into()),
225                ..TokenUsage::default()
226            }),
227            "cached tokens count toward the prompt"
228        );
229        assert_eq!(responses.last(), Some(&LlmResponse::done_with_stop_reason(StopReason::EndTurn)));
230    }
231
232    #[tokio::test]
233    async fn test_metadata_without_cache_fields() {
234        let usage = BedrockTokenUsage::builder().input_tokens(10).output_tokens(5).total_tokens(15).build().unwrap();
235
236        let responses =
237            collect_responses(bedrock_stream().message_stop(BedrockStopReason::EndTurn).metadata(usage).build()).await;
238
239        assert_eq!(responses, llm_response().usage(10, 5).build_with_stop_reason(StopReason::EndTurn));
240    }
241
242    #[tokio::test]
243    async fn test_stop_reasons_map_to_llm_stop_reasons() {
244        for (bedrock_stop_reason, stop_reason) in [
245            (BedrockStopReason::EndTurn, StopReason::EndTurn),
246            (BedrockStopReason::StopSequence, StopReason::EndTurn),
247            (BedrockStopReason::ToolUse, StopReason::ToolCalls),
248            (BedrockStopReason::MaxTokens, StopReason::Length),
249            (BedrockStopReason::ModelContextWindowExceeded, StopReason::Length),
250            (BedrockStopReason::ContentFiltered, StopReason::ContentFilter),
251            (BedrockStopReason::GuardrailIntervened, StopReason::ContentFilter),
252        ] {
253            let responses = collect_responses(bedrock_stream().message_stop(bedrock_stop_reason).build()).await;
254
255            assert_eq!(responses, llm_response().build_with_stop_reason(stop_reason));
256        }
257    }
258
259    async fn collect_responses(events: Vec<ConverseStreamOutput>) -> Vec<LlmResponse> {
260        process_events(events).await.into_iter().map(Result::unwrap).collect()
261    }
262
263    async fn process_events(events: Vec<ConverseStreamOutput>) -> Vec<crate::Result<LlmResponse>> {
264        process_bedrock_stream(stream::iter(events.into_iter().map(Ok))).collect().await
265    }
266
267    fn bedrock_stream() -> BedrockStreamBuilder {
268        BedrockStreamBuilder::default().push(ConverseStreamOutput::MessageStart(
269            MessageStartEvent::builder().role(ConversationRole::Assistant).build().unwrap(),
270        ))
271    }
272
273    #[derive(Default)]
274    struct BedrockStreamBuilder {
275        events: Vec<ConverseStreamOutput>,
276    }
277
278    impl BedrockStreamBuilder {
279        fn text(self, index: i32, chunks: &[&str]) -> Self {
280            chunks
281                .iter()
282                .fold(self, |builder, chunk| builder.delta(index, ContentBlockDelta::Text((*chunk).to_string())))
283                .block_stop(index)
284        }
285
286        fn reasoning(self, index: i32, chunks: &[&str]) -> Self {
287            chunks
288                .iter()
289                .fold(self, |builder, chunk| {
290                    builder.delta(
291                        index,
292                        ContentBlockDelta::ReasoningContent(ReasoningContentBlockDelta::Text((*chunk).to_string())),
293                    )
294                })
295                .block_stop(index)
296        }
297
298        fn tool_call(self, index: i32, id: &str, name: &str, argument_deltas: &[&str]) -> Self {
299            argument_deltas
300                .iter()
301                .fold(self.tool_start(index, id, name), |builder, delta| builder.tool_delta(index, delta))
302                .block_stop(index)
303        }
304
305        fn tool_start(self, index: i32, id: &str, name: &str) -> Self {
306            let tool = ToolUseBlockStart::builder().tool_use_id(id).name(name).build().unwrap();
307            self.push(ConverseStreamOutput::ContentBlockStart(
308                ContentBlockStartEvent::builder()
309                    .content_block_index(index)
310                    .start(ContentBlockStart::ToolUse(tool))
311                    .build()
312                    .unwrap(),
313            ))
314        }
315
316        fn tool_delta(self, index: i32, input: &str) -> Self {
317            self.delta(index, ContentBlockDelta::ToolUse(ToolUseBlockDelta::builder().input(input).build().unwrap()))
318        }
319
320        fn message_stop(self, stop_reason: BedrockStopReason) -> Self {
321            self.push(ConverseStreamOutput::MessageStop(
322                MessageStopEvent::builder().stop_reason(stop_reason).build().unwrap(),
323            ))
324        }
325
326        fn metadata(self, usage: BedrockTokenUsage) -> Self {
327            self.push(ConverseStreamOutput::Metadata(ConverseStreamMetadataEvent::builder().usage(usage).build()))
328        }
329
330        fn delta(self, index: i32, delta: ContentBlockDelta) -> Self {
331            self.push(ConverseStreamOutput::ContentBlockDelta(
332                ContentBlockDeltaEvent::builder().content_block_index(index).delta(delta).build().unwrap(),
333            ))
334        }
335
336        fn block_stop(self, index: i32) -> Self {
337            self.push(ConverseStreamOutput::ContentBlockStop(
338                ContentBlockStopEvent::builder().content_block_index(index).build().unwrap(),
339            ))
340        }
341
342        fn push(mut self, event: ConverseStreamOutput) -> Self {
343            self.events.push(event);
344            self
345        }
346
347        fn build(self) -> Vec<ConverseStreamOutput> {
348            self.events
349        }
350    }
351}