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 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 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 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}