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, StopReason as BedrockStopReason,
6    TokenUsage as BedrockTokenUsage,
7};
8use aws_smithy_types::event_stream::RawMessage;
9use futures::Stream;
10use std::collections::HashMap;
11use tracing::{debug, error, info, warn};
12
13use crate::{LlmError, LlmResponse, ProviderError, StopReason, TokenUsage, Tokens, ToolCallRequest};
14
15impl From<&BedrockTokenUsage> for TokenUsage {
16    fn from(usage: &BedrockTokenUsage) -> Self {
17        let cache_read = usage.cache_read_input_tokens().and_then(|v| u32::try_from(v).ok()).map(Tokens::from);
18        let cache_creation = usage.cache_write_input_tokens().and_then(|v| u32::try_from(v).ok()).map(Tokens::from);
19        // Bedrock's input_tokens excludes cached tokens; TokenUsage counts the whole prompt.
20        let cached = cache_read.unwrap_or_default() + cache_creation.unwrap_or_default();
21        TokenUsage {
22            input_tokens: Tokens::from(u32::try_from(usage.input_tokens).unwrap_or(0)) + cached,
23            output_tokens: u32::try_from(usage.output_tokens).unwrap_or(0).into(),
24            cache_read_tokens: cache_read,
25            cache_creation_tokens: cache_creation,
26            ..TokenUsage::default()
27        }
28    }
29}
30
31struct PendingToolCall {
32    id: String,
33    name: String,
34    args: String,
35}
36
37enum StreamEvent {
38    Emit(LlmResponse),
39    Stop(StopReason),
40    Skip,
41}
42
43pub fn process_bedrock_stream(
44    mut receiver: EventReceiver<ConverseStreamOutput, ConverseStreamOutputError>,
45) -> impl Stream<Item = crate::Result<LlmResponse>> + Send {
46    async_stream::stream! {
47        let message_id = uuid::Uuid::new_v4().to_string();
48        yield Ok(LlmResponse::Start { message_id });
49
50        let mut active_tool_calls: HashMap<i32, PendingToolCall> = HashMap::new();
51        let mut last_stop_reason: Option<StopReason> = None;
52
53        loop {
54            match receiver.recv().await {
55                Ok(Some(event)) => {
56                    match process_stream_event(&event, &mut active_tool_calls) {
57                        StreamEvent::Emit(resp) => yield Ok(resp),
58                        StreamEvent::Stop(sr) => last_stop_reason = Some(sr),
59                        StreamEvent::Skip => {}
60                    }
61                }
62                Ok(None) => {
63                    debug!("Bedrock stream ended (recv returned None)");
64                    break;
65                }
66                Err(e) => {
67                    error!("Bedrock stream recv error: {e}");
68                    yield Err(LlmError::from(e));
69                    return;
70                }
71            }
72        }
73
74        // Emit any remaining tool calls that weren't completed via ContentBlockStop
75        for (_index, tc) in active_tool_calls {
76            let tool_call = ToolCallRequest {
77                id: tc.id,
78                name: tc.name,
79                arguments: tc.args,
80            };
81            yield Ok(LlmResponse::ToolRequestComplete { tool_call });
82        }
83
84        yield Ok(LlmResponse::Done {
85            stop_reason: last_stop_reason,
86        });
87    }
88}
89
90fn process_stream_event(
91    event: &ConverseStreamOutput,
92    active_tool_calls: &mut HashMap<i32, PendingToolCall>,
93) -> StreamEvent {
94    match event {
95        ConverseStreamOutput::MessageStart(_) => {
96            info!("Bedrock message started");
97            StreamEvent::Skip
98        }
99        ConverseStreamOutput::ContentBlockStart(start_event) => {
100            handle_content_block_start(start_event, active_tool_calls)
101        }
102        ConverseStreamOutput::ContentBlockDelta(delta_event) => {
103            handle_content_block_delta(delta_event, active_tool_calls)
104        }
105        ConverseStreamOutput::ContentBlockStop(stop_event) => {
106            handle_content_block_stop(stop_event.content_block_index(), active_tool_calls)
107        }
108        ConverseStreamOutput::MessageStop(stop_event) => {
109            let stop_reason = map_bedrock_stop_reason(&stop_event.stop_reason);
110            info!("Bedrock message stopped: {stop_reason:?}");
111            StreamEvent::Stop(stop_reason)
112        }
113        ConverseStreamOutput::Metadata(metadata_event) => metadata_event
114            .usage()
115            .map_or(StreamEvent::Skip, |usage| StreamEvent::Emit(LlmResponse::Usage { tokens: usage.into() })),
116        other => {
117            warn!("Unhandled Bedrock stream event: {other:?}");
118            StreamEvent::Skip
119        }
120    }
121}
122
123fn handle_content_block_start(
124    event: &aws_sdk_bedrockruntime::types::ContentBlockStartEvent,
125    active_tool_calls: &mut HashMap<i32, PendingToolCall>,
126) -> StreamEvent {
127    let index = event.content_block_index();
128
129    if let Some(ContentBlockStart::ToolUse(tool_start)) = event.start() {
130        let id = tool_start.tool_use_id().to_string();
131        let name = tool_start.name().to_string();
132        debug!("Bedrock tool use started: {name} ({id})");
133        active_tool_calls.insert(index, PendingToolCall { id: id.clone(), name: name.clone(), args: String::new() });
134        StreamEvent::Emit(LlmResponse::ToolRequestStart { id, name })
135    } else {
136        debug!("Content block started at index {index}");
137        StreamEvent::Skip
138    }
139}
140
141fn handle_content_block_delta(
142    event: &aws_sdk_bedrockruntime::types::ContentBlockDeltaEvent,
143    active_tool_calls: &mut HashMap<i32, PendingToolCall>,
144) -> StreamEvent {
145    let index = event.content_block_index();
146
147    let Some(delta) = event.delta() else {
148        return StreamEvent::Skip;
149    };
150
151    match delta {
152        ContentBlockDelta::Text(text) if !text.is_empty() => {
153            StreamEvent::Emit(LlmResponse::Text { chunk: text.clone() })
154        }
155        ContentBlockDelta::ToolUse(tool_delta) => {
156            let input = tool_delta.input();
157            if input.is_empty() {
158                return StreamEvent::Skip;
159            }
160
161            if let Some(tc) = active_tool_calls.get_mut(&index) {
162                tc.args.push_str(input);
163                StreamEvent::Emit(LlmResponse::ToolRequestArg { id: tc.id.clone(), chunk: input.to_string() })
164            } else {
165                warn!("Received tool input delta for unknown content block index: {index}");
166                StreamEvent::Skip
167            }
168        }
169        ContentBlockDelta::ReasoningContent(reasoning) => {
170            if let Ok(text) = reasoning.as_text()
171                && !text.is_empty()
172            {
173                return StreamEvent::Emit(LlmResponse::Reasoning { chunk: text.clone() });
174            }
175            StreamEvent::Skip
176        }
177        _ => {
178            debug!("Unhandled content block delta type");
179            StreamEvent::Skip
180        }
181    }
182}
183
184fn handle_content_block_stop(index: i32, active_tool_calls: &mut HashMap<i32, PendingToolCall>) -> StreamEvent {
185    if let Some(tc) = active_tool_calls.remove(&index) {
186        let tool_call = ToolCallRequest { id: tc.id, name: tc.name, arguments: tc.args };
187        StreamEvent::Emit(LlmResponse::ToolRequestComplete { tool_call })
188    } else {
189        debug!("Content block stopped at index {index}");
190        StreamEvent::Skip
191    }
192}
193
194impl From<SdkError<ConverseStreamOutputError, RawMessage>> for LlmError {
195    fn from(e: SdkError<ConverseStreamOutputError, RawMessage>) -> Self {
196        let message = format!("Bedrock stream error: {e}");
197        let provider = match e {
198            SdkError::ServiceError(svc) => {
199                let inner = svc.err();
200                if inner.is_throttling_exception() {
201                    ProviderError::rate_limit(message)
202                } else if inner.is_service_unavailable_exception()
203                    || inner.is_internal_server_exception()
204                    || inner.is_model_stream_error_exception()
205                {
206                    ProviderError::stream_interrupted(message)
207                } else {
208                    ProviderError::api(message)
209                }
210            }
211            _ => ProviderError::stream_interrupted(message),
212        };
213        Self::from(provider)
214    }
215}
216
217fn map_bedrock_stop_reason(reason: &BedrockStopReason) -> StopReason {
218    match reason {
219        BedrockStopReason::EndTurn | BedrockStopReason::StopSequence => StopReason::EndTurn,
220        BedrockStopReason::ToolUse => StopReason::ToolCalls,
221        BedrockStopReason::MaxTokens | BedrockStopReason::ModelContextWindowExceeded => StopReason::Length,
222        BedrockStopReason::ContentFiltered | BedrockStopReason::GuardrailIntervened => StopReason::ContentFilter,
223        other => StopReason::Unknown(format!("{other:?}")),
224    }
225}
226
227#[cfg(test)]
228mod tests {
229    use super::*;
230
231    #[test]
232    fn test_map_stop_reason_end_turn() {
233        assert_eq!(map_bedrock_stop_reason(&BedrockStopReason::EndTurn), StopReason::EndTurn);
234    }
235
236    #[test]
237    fn test_map_stop_reason_stop_sequence() {
238        assert_eq!(map_bedrock_stop_reason(&BedrockStopReason::StopSequence), StopReason::EndTurn);
239    }
240
241    #[test]
242    fn test_map_stop_reason_tool_use() {
243        assert_eq!(map_bedrock_stop_reason(&BedrockStopReason::ToolUse), StopReason::ToolCalls);
244    }
245
246    #[test]
247    fn test_map_stop_reason_max_tokens() {
248        assert_eq!(map_bedrock_stop_reason(&BedrockStopReason::MaxTokens), StopReason::Length);
249    }
250
251    #[test]
252    fn test_map_stop_reason_context_window_exceeded() {
253        assert_eq!(map_bedrock_stop_reason(&BedrockStopReason::ModelContextWindowExceeded), StopReason::Length);
254    }
255
256    #[test]
257    fn test_map_stop_reason_content_filtered() {
258        assert_eq!(map_bedrock_stop_reason(&BedrockStopReason::ContentFiltered), StopReason::ContentFilter);
259    }
260
261    #[test]
262    fn test_map_stop_reason_guardrail() {
263        assert_eq!(map_bedrock_stop_reason(&BedrockStopReason::GuardrailIntervened), StopReason::ContentFilter);
264    }
265
266    #[test]
267    fn test_handle_content_block_start_tool_use() {
268        let mut active = HashMap::new();
269        let tool_start = aws_sdk_bedrockruntime::types::ToolUseBlockStart::builder()
270            .tool_use_id("tool_123")
271            .name("search")
272            .build()
273            .unwrap();
274
275        let event = aws_sdk_bedrockruntime::types::ContentBlockStartEvent::builder()
276            .content_block_index(0)
277            .start(ContentBlockStart::ToolUse(tool_start))
278            .build()
279            .unwrap();
280
281        let result = handle_content_block_start(&event, &mut active);
282        assert!(
283            matches!(&result, StreamEvent::Emit(LlmResponse::ToolRequestStart { id, name }) if id == "tool_123" && name == "search")
284        );
285        assert!(active.contains_key(&0));
286    }
287
288    #[test]
289    fn test_handle_content_block_delta_text() {
290        let mut active = HashMap::new();
291        let delta = aws_sdk_bedrockruntime::types::ContentBlockDeltaEvent::builder()
292            .content_block_index(0)
293            .delta(ContentBlockDelta::Text("Hello".to_string()))
294            .build()
295            .unwrap();
296
297        let result = handle_content_block_delta(&delta, &mut active);
298        assert!(matches!(&result, StreamEvent::Emit(LlmResponse::Text { chunk }) if chunk == "Hello"));
299    }
300
301    #[test]
302    fn test_handle_content_block_delta_tool_input() {
303        let mut active = HashMap::new();
304        active
305            .insert(0, PendingToolCall { id: "tool_123".to_string(), name: "search".to_string(), args: String::new() });
306
307        let tool_delta =
308            aws_sdk_bedrockruntime::types::ToolUseBlockDelta::builder().input(r#"{"query":"test"}"#).build().unwrap();
309
310        let delta = aws_sdk_bedrockruntime::types::ContentBlockDeltaEvent::builder()
311            .content_block_index(0)
312            .delta(ContentBlockDelta::ToolUse(tool_delta))
313            .build()
314            .unwrap();
315
316        let result = handle_content_block_delta(&delta, &mut active);
317        assert!(
318            matches!(&result, StreamEvent::Emit(LlmResponse::ToolRequestArg { id, chunk }) if id == "tool_123" && chunk == r#"{"query":"test"}"#)
319        );
320
321        // Verify accumulated args
322        assert_eq!(active.get(&0).unwrap().args, r#"{"query":"test"}"#);
323    }
324
325    #[test]
326    fn test_handle_content_block_stop_completes_tool() {
327        let mut active = HashMap::new();
328        active.insert(
329            0,
330            PendingToolCall {
331                id: "tool_123".to_string(),
332                name: "search".to_string(),
333                args: r#"{"query":"test"}"#.to_string(),
334            },
335        );
336
337        let result = handle_content_block_stop(0, &mut active);
338        assert!(matches!(&result, StreamEvent::Emit(LlmResponse::ToolRequestComplete { tool_call })
339            if tool_call.id == "tool_123"
340            && tool_call.name == "search"
341            && tool_call.arguments == r#"{"query":"test"}"#
342        ));
343        assert!(active.is_empty());
344    }
345
346    #[test]
347    fn test_handle_content_block_stop_no_tool() {
348        let mut active = HashMap::new();
349        let result = handle_content_block_stop(0, &mut active);
350        assert!(matches!(result, StreamEvent::Skip));
351    }
352
353    #[test]
354    fn test_metadata_event_emits_cache_read_and_creation() {
355        let usage = aws_sdk_bedrockruntime::types::TokenUsage::builder()
356            .input_tokens(100)
357            .output_tokens(50)
358            .total_tokens(150)
359            .cache_read_input_tokens(40)
360            .cache_write_input_tokens(20)
361            .build()
362            .unwrap();
363
364        let metadata = aws_sdk_bedrockruntime::types::ConverseStreamMetadataEvent::builder().usage(usage).build();
365
366        let event = ConverseStreamOutput::Metadata(metadata);
367        let mut active = HashMap::new();
368        let result = process_stream_event(&event, &mut active);
369
370        match result {
371            StreamEvent::Emit(LlmResponse::Usage { tokens: sample }) => {
372                assert_eq!(sample.input_tokens.get(), 160, "cached tokens count toward the prompt");
373                assert_eq!(sample.output_tokens.get(), 50);
374                assert_eq!(sample.cache_read_tokens.map(crate::Tokens::get), Some(40));
375                assert_eq!(sample.cache_creation_tokens.map(crate::Tokens::get), Some(20));
376            }
377            _ => panic!("expected Emit(Usage{{..}})"),
378        }
379    }
380
381    #[test]
382    fn test_metadata_event_without_cache_fields() {
383        let usage = aws_sdk_bedrockruntime::types::TokenUsage::builder()
384            .input_tokens(10)
385            .output_tokens(5)
386            .total_tokens(15)
387            .build()
388            .unwrap();
389
390        let metadata = aws_sdk_bedrockruntime::types::ConverseStreamMetadataEvent::builder().usage(usage).build();
391
392        let event = ConverseStreamOutput::Metadata(metadata);
393        let mut active = HashMap::new();
394        let result = process_stream_event(&event, &mut active);
395
396        match result {
397            StreamEvent::Emit(LlmResponse::Usage { tokens: sample }) => {
398                assert_eq!(sample.cache_read_tokens, None);
399                assert_eq!(sample.cache_creation_tokens, None);
400            }
401            _ => panic!("expected Emit(Usage{{..}})"),
402        }
403    }
404
405    #[test]
406    fn test_handle_content_block_delta_empty_text() {
407        let mut active = HashMap::new();
408        let delta = aws_sdk_bedrockruntime::types::ContentBlockDeltaEvent::builder()
409            .content_block_index(0)
410            .delta(ContentBlockDelta::Text(String::new()))
411            .build()
412            .unwrap();
413
414        let result = handle_content_block_delta(&delta, &mut active);
415        assert!(matches!(result, StreamEvent::Skip));
416    }
417}