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