From 76c8028c26a2f88a3a1f5249aaf169432ba543b6 Mon Sep 17 00:00:00 2001 From: Simon Ho Date: Tue, 16 Jun 2026 19:20:05 +0100 Subject: [PATCH] fix(bedrock): implement real streaming via ConverseStream API (#9579) Signed-off-by: Simon Ho --- crates/goose/src/providers/bedrock.rs | 826 ++++++++++++++++++++++++-- 1 file changed, 787 insertions(+), 39 deletions(-) diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index d364937f3..664b1aa96 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -6,15 +6,19 @@ use crate::conversation::message::Message; use crate::model::ModelConfig; use crate::providers::utils::RequestLog; use anyhow::Result; +use async_stream::try_stream; use async_trait::async_trait; use aws_sdk_bedrockruntime::config::ProvideCredentials; use aws_sdk_bedrockruntime::operation::converse::ConverseError; +use aws_sdk_bedrockruntime::operation::converse_stream::ConverseStreamError; +use aws_sdk_bedrockruntime::types::error::ConverseStreamOutputError; use aws_sdk_bedrockruntime::{types as bedrock, Client}; +use base64::Engine; use futures::future::BoxFuture; -use goose_providers::conversation::token_usage::ProviderUsage; +use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; use goose_providers::errors::ProviderError; use reqwest::header::HeaderValue; -use rmcp::model::Tool; +use rmcp::model::{object, CallToolRequestParams, ErrorCode, ErrorData, Tool}; use serde_json::Value; use smithy_transport_reqwest::ReqwestHttpClient; @@ -53,6 +57,13 @@ pub struct BedrockProvider { name: String, } +/// Request inputs shared by the `Converse` and `ConverseStream` APIs. +struct ConverseRequestParts { + system_blocks: Vec, + messages: Vec, + tool_config: Option, +} + impl BedrockProvider { pub async fn from_env(model: ModelConfig) -> Result { let config = crate::config::Config::global(); @@ -199,15 +210,16 @@ impl BedrockProvider { enabled && self.model.model_name.contains("anthropic.claude") } - async fn converse( + /// Build the request inputs shared by [`Self::converse`] and + /// [`Self::converse_stream`]: system blocks (with optional cache point), + /// converted messages (with optional trailing-message cache point), and + /// the tool configuration. + fn build_request_parts( &self, - session_id: Option<&str>, system: &str, messages: &[Message], tools: &[Tool], - ) -> Result<(bedrock::Message, Option), ProviderError> { - let model_name = &self.model.model_name; - + ) -> Result { let enable_caching = self.should_enable_caching(); let system_blocks = if enable_caching { @@ -235,23 +247,45 @@ impl BedrockProvider { let last_idx = visible_messages.len().saturating_sub(1); + let bedrock_messages = visible_messages + .iter() + .enumerate() + .map(|(idx, m)| to_bedrock_message_with_caching(m, enable_caching && idx == last_idx)) + .collect::>>()?; + + let tool_config = if tools.is_empty() { + None + } else { + Some(to_bedrock_tool_config(tools)?) + }; + + Ok(ConverseRequestParts { + system_blocks, + messages: bedrock_messages, + tool_config, + }) + } + + async fn converse( + &self, + session_id: Option<&str>, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result<(bedrock::Message, Option), ProviderError> { + let model_name = &self.model.model_name; + + let parts = self.build_request_parts(system, messages, tools)?; + let mut request = self .client .converse() - .set_system(Some(system_blocks)) + .set_system(Some(parts.system_blocks)) .model_id(model_name.to_string()) - .set_messages(Some( - visible_messages - .iter() - .enumerate() - .map(|(idx, m)| { - to_bedrock_message_with_caching(m, enable_caching && idx == last_idx) - }) - .collect::>()?, - )); + .set_messages(Some(parts.messages)); - if !tools.is_empty() { - request = request.tool_config(to_bedrock_tool_config(tools)?); + if let Some(tool_config) = parts.tool_config { + request = request.tool_config(tool_config); } let mut request = request.customize(); @@ -307,6 +341,268 @@ impl BedrockProvider { )), } } + + /// Escape hatch: `BEDROCK_DISABLE_STREAMING=true` restores the previous + /// blocking `Converse` behaviour in case a model or region misbehaves + /// with `ConverseStream`. + fn streaming_disabled(&self) -> bool { + let config = crate::config::Config::global(); + config + .get_param::("BEDROCK_DISABLE_STREAMING") + .unwrap_or(false) + } + + /// Streaming variant of [`Self::converse`]. Builds an identical request + /// but calls the AWS `ConverseStream` API, returning the raw event + /// receiver so [`Provider::stream`] can forward deltas incrementally. + async fn converse_stream( + &self, + session_id: Option<&str>, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result< + aws_sdk_bedrockruntime::operation::converse_stream::ConverseStreamOutput, + ProviderError, + > { + let model_name = &self.model.model_name; + + let parts = self.build_request_parts(system, messages, tools)?; + + let mut request = self + .client + .converse_stream() + .set_system(Some(parts.system_blocks)) + .model_id(model_name.to_string()) + .set_messages(Some(parts.messages)); + + if let Some(tool_config) = parts.tool_config { + request = request.tool_config(tool_config); + } + + let mut request = request.customize(); + + if let Some(session_id) = session_id.filter(|id| !id.is_empty()) { + let session_id = session_id.to_string(); + request = request.mutate_request(move |req| { + if let Ok(value) = HeaderValue::from_str(&session_id) { + req.headers_mut().insert(SESSION_ID_HEADER, value); + } + }); + } + + request + .send() + .await + .map_err(|err| match err.into_service_error() { + ConverseStreamError::ThrottlingException(throttle_err) => { + ProviderError::RateLimitExceeded { + details: format!("Bedrock throttling error: {:?}", throttle_err), + retry_delay: None, + } + } + ConverseStreamError::AccessDeniedException(err) => { + ProviderError::Authentication(format!("Failed to call Bedrock: {:?}", err)) + } + ConverseStreamError::ValidationException(err) + if { + let msg = err.message().unwrap_or_default(); + msg.contains("Input is too long for requested model.") + || msg.contains("prompt is too long") + } => + { + ProviderError::ContextLengthExceeded(format!( + "Failed to call Bedrock: {:?}", + err + )) + } + ConverseStreamError::ModelErrorException(err) => { + ProviderError::ExecutionError(format!("Failed to call Bedrock: {:?}", err)) + } + err => ProviderError::ServerError(format!("Failed to call Bedrock: {:?}", err)), + }) + } + + /// Pre-ConverseStream behaviour: blocking `Converse` call wrapped in a + /// single-item stream. Kept as the `BEDROCK_DISABLE_STREAMING=true` + /// escape hatch. + async fn stream_via_converse( + &self, + session_id: Option<&str>, + system: &str, + messages: &[Message], + tools: &[Tool], + model_name: &str, + ) -> Result { + let (bedrock_message, bedrock_usage) = self + .with_retry(|| self.converse(session_id, system, messages, tools)) + .await?; + + let usage = bedrock_usage + .as_ref() + .map(from_bedrock_usage) + .unwrap_or_default(); + + let message = from_bedrock_message(&bedrock_message)?; + + // Add debug trace with input context + let debug_payload = serde_json::json!({ + "system": system, + "messages": messages, + "tools": tools + }); + let mut log = RequestLog::start(&self.model, &debug_payload)?; + log.write( + &serde_json::to_value(&message).unwrap_or_default(), + Some(&usage), + )?; + + let provider_usage = ProviderUsage::new(model_name.to_string(), usage); + Ok(super::base::stream_from_single_message( + message, + provider_usage, + )) + } +} + +/// Accumulation state for in-flight content blocks while consuming a +/// `ConverseStream` response. Tool inputs and reasoning content arrive as +/// fragments that only become a complete [`Message`] at `ContentBlockStop`. +#[derive(Default)] +struct StreamBlockState { + /// content_block_index -> (tool_use_id, tool_name, accumulated input JSON) + tool_blocks: HashMap, + /// content_block_index -> (accumulated reasoning text, accumulated signature) + reasoning_blocks: HashMap, + /// content_block_index -> accumulated redacted (encrypted) reasoning bytes + redacted_blocks: HashMap>, +} + +/// Convert a single `ConverseStream` event into zero or more [`Message`]s +/// ready to be yielded, plus token usage when the event carries it. +/// +/// Mirrors the delta-yield contract of +/// `formats::anthropic::response_to_streaming_message`: text deltas yield +/// immediately (token-level chunks); tool-use inputs and reasoning blocks +/// accumulate in `state` until their `ContentBlockStop`. +fn process_stream_event( + event: bedrock::ConverseStreamOutput, + state: &mut StreamBlockState, + message_id: &str, +) -> (Vec, Option) { + let mut messages = Vec::new(); + let mut usage = None; + + match event { + bedrock::ConverseStreamOutput::ContentBlockStart(ev) => { + if let Some(bedrock::ContentBlockStart::ToolUse(tu)) = ev.start { + state.tool_blocks.insert( + ev.content_block_index, + (tu.tool_use_id, tu.name, String::new()), + ); + } + } + bedrock::ConverseStreamOutput::ContentBlockDelta(ev) => match ev.delta { + Some(bedrock::ContentBlockDelta::Text(text)) => { + if !text.is_empty() { + messages.push(Message::assistant().with_text(text).with_id(message_id)); + } + } + Some(bedrock::ContentBlockDelta::ToolUse(tu)) => { + if let Some(entry) = state.tool_blocks.get_mut(&ev.content_block_index) { + entry.2.push_str(&tu.input); + } + } + Some(bedrock::ContentBlockDelta::ReasoningContent(rc)) => match rc { + bedrock::ReasoningContentBlockDelta::Text(t) => { + state + .reasoning_blocks + .entry(ev.content_block_index) + .or_default() + .0 + .push_str(&t); + } + bedrock::ReasoningContentBlockDelta::Signature(s) => { + state + .reasoning_blocks + .entry(ev.content_block_index) + .or_default() + .1 + .push_str(&s); + } + bedrock::ReasoningContentBlockDelta::RedactedContent(blob) => { + state + .redacted_blocks + .entry(ev.content_block_index) + .or_default() + .extend_from_slice(blob.as_ref()); + } + _ => {} + }, + _ => {} + }, + bedrock::ConverseStreamOutput::ContentBlockStop(ev) => { + let idx = ev.content_block_index; + if let Some((text, signature)) = state.reasoning_blocks.remove(&idx) { + if !text.is_empty() { + messages.push( + Message::assistant() + .with_thinking(text, signature) + .with_id(message_id), + ); + } + } + if let Some(bytes) = state.redacted_blocks.remove(&idx) { + if !bytes.is_empty() { + // Same base64 encoding as the non-streaming path + // (formats::bedrock::from_bedrock_reasoning_content_block) + // so redacted thinking round-trips back to Bedrock intact. + let encoded = base64::prelude::BASE64_STANDARD.encode(&bytes); + messages.push( + Message::assistant() + .with_redacted_thinking(encoded) + .with_id(message_id), + ); + } + } + if let Some((id, name, input_json)) = state.tool_blocks.remove(&idx) { + // Parse the accumulated tool input. On failure, yield an + // error tool request (not a stream error) so the agent can + // report it back to the model — same behaviour as the + // Anthropic provider. + let tool_call = if input_json.trim().is_empty() { + Ok(CallToolRequestParams::new(name) + .with_arguments(object(serde_json::json!({})))) + } else { + match serde_json::from_str::(&input_json) { + Ok(parsed) => { + Ok(CallToolRequestParams::new(name).with_arguments(object(parsed))) + } + Err(_) => Err(ErrorData::new( + ErrorCode::INVALID_PARAMS, + format!("Could not parse tool arguments: {}", input_json), + None, + )), + } + }; + messages.push( + Message::assistant() + .with_tool_request(id, tool_call) + .with_id(message_id), + ); + } + } + bedrock::ConverseStreamOutput::Metadata(ev) => { + if let Some(u) = ev.usage { + usage = Some(from_bedrock_usage(&u)); + } + } + // MessageStart / MessageStop / unknown variants carry no content + // that needs forwarding. + _ => {} + } + + (messages, usage) } impl ProviderDef for BedrockProvider { @@ -316,7 +612,7 @@ impl ProviderDef for BedrockProvider { ProviderMetadata::new( BEDROCK_PROVIDER_NAME, "Amazon Bedrock", - "Run models through Amazon Bedrock. Supports AWS SSO profiles - run 'aws sso login --profile ' before using. Configure with AWS_PROFILE and AWS_REGION, use environment variables/credentials, or use AWS_BEARER_TOKEN_BEDROCK for bearer token authentication. Region is required for bearer token auth (can be set via AWS_REGION, AWS_DEFAULT_REGION, or AWS profile). Prompt caching can be enabled for Anthropic Claude models by setting BEDROCK_ENABLE_CACHING=true.", + "Run models through Amazon Bedrock. Supports AWS SSO profiles - run 'aws sso login --profile ' before using. Configure with AWS_PROFILE and AWS_REGION, use environment variables/credentials, or use AWS_BEARER_TOKEN_BEDROCK for bearer token authentication. Region is required for bearer token auth (can be set via AWS_REGION, AWS_DEFAULT_REGION, or AWS profile). Prompt caching can be enabled for Anthropic Claude models by setting BEDROCK_ENABLE_CACHING=true. Responses stream via the ConverseStream API; set BEDROCK_DISABLE_STREAMING=true to fall back to blocking Converse calls.", BEDROCK_DEFAULT_MODEL, BEDROCK_KNOWN_MODELS.to_vec(), BEDROCK_DOC_LINK, @@ -325,6 +621,13 @@ impl ProviderDef for BedrockProvider { ConfigKey::new("AWS_REGION", true, false, Some("us-east-1"), true), ConfigKey::new("AWS_BEARER_TOKEN_BEDROCK", false, true, None, true), ConfigKey::new("BEDROCK_ENABLE_CACHING", false, false, Some("false"), false), + ConfigKey::new( + "BEDROCK_DISABLE_STREAMING", + false, + false, + Some("false"), + false, + ), ], ) } @@ -370,34 +673,111 @@ impl Provider for BedrockProvider { }; let model_name = model_config.model_name.clone(); - let (bedrock_message, bedrock_usage) = self - .with_retry(|| self.converse(session_id, system, messages, tools)) + // Escape hatch: restore the previous blocking-Converse behaviour. + if self.streaming_disabled() { + return self + .stream_via_converse(session_id, system, messages, tools, &model_name) + .await; + } + + // Open the AWS ConverseStream event stream. Retry wraps the request + // setup only — mid-stream errors are surfaced, not retried (matching + // the Anthropic provider's behaviour). + let response = self + .with_retry(|| self.converse_stream(session_id, system, messages, tools)) .await?; - let usage = bedrock_usage - .as_ref() - .map(from_bedrock_usage) - .unwrap_or_default(); - - let message = from_bedrock_message(&bedrock_message)?; - - // Add debug trace with input context + // Debug trace with input context; the streamed text is written once + // the stream completes. let debug_payload = serde_json::json!({ "system": system, "messages": messages, "tools": tools }); let mut log = RequestLog::start(&self.model, &debug_payload)?; - log.write( - &serde_json::to_value(&message).unwrap_or_default(), - Some(&usage), - )?; - let provider_usage = ProviderUsage::new(model_name.to_string(), usage); - Ok(super::base::stream_from_single_message( - message, - provider_usage, - )) + let mut event_stream = response.stream; + + Ok(Box::pin(try_stream! { + let mut state = StreamBlockState::default(); + // One id for the whole assistant turn so consumers + // (Conversation::push) can coalesce consecutive deltas into a + // single message — mirrors the Anthropic provider, which stamps + // the API-provided message id on every chunk. Bedrock's + // MessageStart event carries no id, so generate one. + let message_id = format!("msg_{}", uuid::Uuid::new_v4()); + let mut full_text = String::new(); + let mut final_usage: Option = None; + + loop { + let event = event_stream.recv().await.map_err(|err| { + // Map Bedrock mid-stream exceptions to specific ProviderError + // variants so the agent's retry / context-length / server-error + // handling kicks in, mirroring the non-streaming Converse error + // mapping. Without this, a mid-stream throttling or + // context-length failure would be flattened to a generic + // RequestFailed and lose its retryable / context semantics. + match err.as_service_error() { + Some(ConverseStreamOutputError::ThrottlingException(e)) => { + ProviderError::RateLimitExceeded { + details: format!("Bedrock streaming throttling error: {:?}", e), + retry_delay: None, + } + } + Some(ConverseStreamOutputError::ValidationException(e)) + if { + let msg = e.message().unwrap_or_default(); + msg.contains("Input is too long for requested model.") + || msg.contains("prompt is too long") + } => + { + ProviderError::ContextLengthExceeded(format!( + "Bedrock streaming validation error: {:?}", + e + )) + } + Some(ConverseStreamOutputError::ServiceUnavailableException(_)) + | Some(ConverseStreamOutputError::InternalServerException(_)) => { + ProviderError::ServerError(format!( + "Bedrock streaming server error: {:?}", + err + )) + } + Some(ConverseStreamOutputError::ModelStreamErrorException(e)) => { + ProviderError::ExecutionError(format!( + "Bedrock model stream error: {:?}", + e + )) + } + _ => ProviderError::RequestFailed(format!( + "Bedrock stream receive error: {:?}", + err + )), + } + })?; + let Some(event) = event else { break }; + + let (messages, usage) = process_stream_event(event, &mut state, &message_id); + if let Some(usage) = usage { + final_usage = Some(ProviderUsage::new(model_name.clone(), usage)); + } + for message in messages { + if let Some(text) = message.content.first().and_then(|c| c.as_text()) { + full_text.push_str(text); + } + yield (Some(message), None); + } + } + + let usage = final_usage.unwrap_or_else(|| { + ProviderUsage::new(model_name.clone(), Usage::default()) + }); + let _ = log.write( + &serde_json::json!({ "streamed_text": full_text }), + Some(&usage.usage), + ); + yield (None, Some(usage)); + })) } } @@ -528,4 +908,372 @@ mod tests { std::env::remove_var("BEDROCK_ENABLE_CACHING"); } + + // ── ConverseStream event processing ────────────────────────────────── + + use crate::conversation::message::MessageContent; + + /// Stand-in for the per-turn message id that `stream()` generates. + const TEST_MESSAGE_ID: &str = "msg_test"; + + fn delta_event(idx: i32, delta: bedrock::ContentBlockDelta) -> bedrock::ConverseStreamOutput { + bedrock::ConverseStreamOutput::ContentBlockDelta( + bedrock::ContentBlockDeltaEvent::builder() + .delta(delta) + .content_block_index(idx) + .build() + .unwrap(), + ) + } + + fn tool_start_event(idx: i32, id: &str, name: &str) -> bedrock::ConverseStreamOutput { + bedrock::ConverseStreamOutput::ContentBlockStart( + bedrock::ContentBlockStartEvent::builder() + .start(bedrock::ContentBlockStart::ToolUse( + bedrock::ToolUseBlockStart::builder() + .tool_use_id(id) + .name(name) + .build() + .unwrap(), + )) + .content_block_index(idx) + .build() + .unwrap(), + ) + } + + fn tool_delta_event(idx: i32, fragment: &str) -> bedrock::ConverseStreamOutput { + delta_event( + idx, + bedrock::ContentBlockDelta::ToolUse( + bedrock::ToolUseBlockDelta::builder() + .input(fragment) + .build() + .unwrap(), + ), + ) + } + + fn stop_event(idx: i32) -> bedrock::ConverseStreamOutput { + bedrock::ConverseStreamOutput::ContentBlockStop( + bedrock::ContentBlockStopEvent::builder() + .content_block_index(idx) + .build() + .unwrap(), + ) + } + + #[test] + fn test_stream_text_delta_yields_immediately() { + let mut state = StreamBlockState::default(); + let (messages, usage) = process_stream_event( + delta_event(0, bedrock::ContentBlockDelta::Text("Hello".to_string())), + &mut state, + TEST_MESSAGE_ID, + ); + assert_eq!(messages.len(), 1, "text delta should yield one message"); + assert_eq!(messages[0].as_concat_text(), "Hello"); + assert!(usage.is_none()); + } + + #[test] + fn test_stream_empty_text_delta_yields_nothing() { + let mut state = StreamBlockState::default(); + let (messages, _) = process_stream_event( + delta_event(0, bedrock::ContentBlockDelta::Text(String::new())), + &mut state, + TEST_MESSAGE_ID, + ); + assert!(messages.is_empty(), "empty text delta should be skipped"); + } + + #[test] + fn test_stream_tool_use_accumulates_until_stop() { + let mut state = StreamBlockState::default(); + + let (messages, _) = process_stream_event( + tool_start_event(1, "tool-1", "file_write"), + &mut state, + TEST_MESSAGE_ID, + ); + assert!(messages.is_empty(), "tool start should not yield"); + + // Input arrives as partial-JSON fragments + for fragment in [r#"{"path": "#, r#""a.txt"}"#] { + let (messages, _) = + process_stream_event(tool_delta_event(1, fragment), &mut state, TEST_MESSAGE_ID); + assert!(messages.is_empty(), "tool input fragments should not yield"); + } + + let (messages, _) = process_stream_event(stop_event(1), &mut state, TEST_MESSAGE_ID); + assert_eq!( + messages.len(), + 1, + "tool stop should yield the complete request" + ); + match &messages[0].content[0] { + MessageContent::ToolRequest(req) => { + assert_eq!(req.id, "tool-1"); + let call = req + .tool_call + .as_ref() + .expect("accumulated JSON should parse"); + assert_eq!(call.name.to_string(), "file_write"); + let args = call.arguments.as_ref().expect("arguments should be set"); + assert_eq!(args.get("path").and_then(|v| v.as_str()), Some("a.txt")); + } + other => panic!("expected ToolRequest, got {:?}", other), + } + } + + #[test] + fn test_stream_tool_use_invalid_json_yields_error_request() { + let mut state = StreamBlockState::default(); + + process_stream_event( + tool_start_event(0, "tool-2", "shell"), + &mut state, + TEST_MESSAGE_ID, + ); + process_stream_event( + tool_delta_event(0, "this is {{{ not json"), + &mut state, + TEST_MESSAGE_ID, + ); + + let (messages, _) = process_stream_event(stop_event(0), &mut state, TEST_MESSAGE_ID); + assert_eq!(messages.len(), 1); + match &messages[0].content[0] { + MessageContent::ToolRequest(req) => { + assert!( + req.tool_call.is_err(), + "unparseable input should yield an error tool request, not a stream failure" + ); + } + other => panic!("expected ToolRequest, got {:?}", other), + } + } + + #[test] + fn test_stream_tool_use_empty_input_yields_empty_args() { + let mut state = StreamBlockState::default(); + + process_stream_event( + tool_start_event(0, "tool-3", "list_files"), + &mut state, + TEST_MESSAGE_ID, + ); + // No input deltas at all — some tools take no arguments. + + let (messages, _) = process_stream_event(stop_event(0), &mut state, TEST_MESSAGE_ID); + assert_eq!(messages.len(), 1); + match &messages[0].content[0] { + MessageContent::ToolRequest(req) => { + let call = req + .tool_call + .as_ref() + .expect("empty input should parse as {}"); + assert_eq!(call.name.to_string(), "list_files"); + } + other => panic!("expected ToolRequest, got {:?}", other), + } + } + + #[test] + fn test_stream_reasoning_accumulates_until_stop() { + let mut state = StreamBlockState::default(); + + for (delta, expect_empty) in [ + ( + bedrock::ReasoningContentBlockDelta::Text("Let me think".to_string()), + true, + ), + ( + bedrock::ReasoningContentBlockDelta::Text(" about this.".to_string()), + true, + ), + ( + bedrock::ReasoningContentBlockDelta::Signature("sig-abc".to_string()), + true, + ), + ] { + let (messages, _) = process_stream_event( + delta_event(0, bedrock::ContentBlockDelta::ReasoningContent(delta)), + &mut state, + TEST_MESSAGE_ID, + ); + assert_eq!( + messages.is_empty(), + expect_empty, + "reasoning deltas accumulate" + ); + } + + let (messages, _) = process_stream_event(stop_event(0), &mut state, TEST_MESSAGE_ID); + assert_eq!( + messages.len(), + 1, + "reasoning stop should yield thinking message" + ); + match &messages[0].content[0] { + MessageContent::Thinking(t) => { + assert_eq!(t.thinking, "Let me think about this."); + assert_eq!(t.signature, "sig-abc"); + } + other => panic!("expected Thinking, got {:?}", other), + } + } + + #[test] + fn test_stream_metadata_returns_usage() { + let mut state = StreamBlockState::default(); + let event = bedrock::ConverseStreamOutput::Metadata( + bedrock::ConverseStreamMetadataEvent::builder() + .usage( + bedrock::TokenUsage::builder() + .input_tokens(100) + .output_tokens(50) + .total_tokens(150) + .build() + .unwrap(), + ) + .build(), + ); + let (messages, usage) = process_stream_event(event, &mut state, TEST_MESSAGE_ID); + assert!(messages.is_empty()); + let usage = usage.expect("metadata event should carry usage"); + assert_eq!(usage.input_tokens, Some(100)); + assert_eq!(usage.output_tokens, Some(50)); + assert_eq!(usage.total_tokens, Some(150)); + } + + #[test] + fn test_stream_interleaved_text_and_tool_blocks() { + // Bedrock interleaves block indices: text at index 0, tool at index 1. + // Text yields immediately even while a tool block is mid-accumulation. + let mut state = StreamBlockState::default(); + + process_stream_event( + tool_start_event(1, "tool-4", "search"), + &mut state, + TEST_MESSAGE_ID, + ); + process_stream_event(tool_delta_event(1, r#"{"q":"#), &mut state, TEST_MESSAGE_ID); + + let (messages, _) = process_stream_event( + delta_event( + 0, + bedrock::ContentBlockDelta::Text("Searching now".to_string()), + ), + &mut state, + TEST_MESSAGE_ID, + ); + assert_eq!( + messages.len(), + 1, + "text should stream while tool accumulates" + ); + assert_eq!(messages[0].as_concat_text(), "Searching now"); + + process_stream_event( + tool_delta_event(1, r#""rust"}"#), + &mut state, + TEST_MESSAGE_ID, + ); + let (messages, _) = process_stream_event(stop_event(1), &mut state, TEST_MESSAGE_ID); + assert_eq!(messages.len(), 1); + match &messages[0].content[0] { + MessageContent::ToolRequest(req) => { + let call = req.tool_call.as_ref().unwrap(); + let args = call.arguments.as_ref().unwrap(); + assert_eq!(args.get("q").and_then(|v| v.as_str()), Some("rust")); + } + other => panic!("expected ToolRequest, got {:?}", other), + } + } + + #[test] + fn test_metadata_includes_disable_streaming_key() { + let meta = BedrockProvider::metadata(); + + let key = meta + .config_keys + .iter() + .find(|k| k.name == "BEDROCK_DISABLE_STREAMING") + .expect("BEDROCK_DISABLE_STREAMING config key should exist"); + assert!( + !key.required, + "BEDROCK_DISABLE_STREAMING should not be required" + ); + assert!( + !key.secret, + "BEDROCK_DISABLE_STREAMING should not be marked as secret" + ); + assert_eq!( + key.default.as_deref(), + Some("false"), + "BEDROCK_DISABLE_STREAMING should default to false (streaming on)" + ); + } + + #[test] + fn test_stream_messages_carry_turn_message_id() { + // Every message from one turn must share the caller-provided id so + // Conversation::push can coalesce consecutive deltas instead of + // persisting one message per token (same contract as the Anthropic + // provider, which stamps the API message id on every chunk). + let mut state = StreamBlockState::default(); + + let (messages, _) = process_stream_event( + delta_event(0, bedrock::ContentBlockDelta::Text("Hello".to_string())), + &mut state, + TEST_MESSAGE_ID, + ); + assert_eq!(messages[0].id.as_deref(), Some(TEST_MESSAGE_ID)); + + process_stream_event( + tool_start_event(1, "tool-9", "shell"), + &mut state, + TEST_MESSAGE_ID, + ); + let (messages, _) = process_stream_event(stop_event(1), &mut state, TEST_MESSAGE_ID); + assert_eq!( + messages[0].id.as_deref(), + Some(TEST_MESSAGE_ID), + "tool requests must carry the same turn id as text deltas" + ); + } + + #[test] + fn test_stream_redacted_reasoning_accumulates_until_stop() { + let mut state = StreamBlockState::default(); + + let raw = b"encrypted-reasoning-bytes"; + let (messages, _) = process_stream_event( + delta_event( + 0, + bedrock::ContentBlockDelta::ReasoningContent( + bedrock::ReasoningContentBlockDelta::RedactedContent( + aws_smithy_types::Blob::new(raw.to_vec()), + ), + ), + ), + &mut state, + TEST_MESSAGE_ID, + ); + assert!(messages.is_empty(), "redacted deltas accumulate until stop"); + + let (messages, _) = process_stream_event(stop_event(0), &mut state, TEST_MESSAGE_ID); + assert_eq!(messages.len(), 1); + match &messages[0].content[0] { + MessageContent::RedactedThinking(redacted) => { + let expected = base64::prelude::BASE64_STANDARD.encode(raw); + assert_eq!( + redacted.data, expected, + "blob must round-trip as base64, matching the non-streaming path" + ); + } + other => panic!("expected RedactedThinking, got {:?}", other), + } + } }