use super::{ChatEvent, ChatProvider, ChatUsage, SamplingParams, ToolDef};
use crate::ChatMessage;
use anyhow::{Context, Result, anyhow};
use async_trait::async_trait;
use aws_config::BehaviorVersion;
use aws_sdk_bedrockruntime::Client as BedrockClient;
use aws_sdk_bedrockruntime::types::ConverseStreamOutput as StreamEvent;
use aws_sdk_bedrockruntime::types::{
ContentBlock, ContentBlockDelta, ConversationRole, InferenceConfiguration, Message,
SystemContentBlock,
};
use tokio::sync::mpsc::Sender;
pub const DEFAULT_BEDROCK_MODEL: &str = "us.anthropic.claude-sonnet-4-6";
pub const ENV_REGION_TRUSTY: &str = "TRUSTY_AWS_REGION";
pub const ENV_REGION_AWS: &str = "AWS_REGION";
pub const DEFAULT_BEDROCK_REGION: &str = "us-east-1";
pub fn resolve_bedrock_region(explicit: Option<&str>) -> String {
if let Some(r) = explicit.filter(|s| !s.is_empty()) {
return r.to_string();
}
for var in [ENV_REGION_TRUSTY, ENV_REGION_AWS] {
let val = std::env::var(var).unwrap_or_default();
if !val.is_empty() {
return val;
}
}
DEFAULT_BEDROCK_REGION.to_string()
}
pub struct BedrockProvider {
client: BedrockClient,
model: String,
region: String,
sampling: SamplingParams,
}
impl BedrockProvider {
pub async fn new(model: impl Into<String>, region: Option<&str>) -> Result<Self> {
let region_str = resolve_bedrock_region(region);
let region_provider = aws_config::meta::region::RegionProviderChain::first_try(
aws_types::region::Region::new(region_str.clone()),
);
let config = aws_config::defaults(BehaviorVersion::latest())
.region(region_provider)
.load()
.await;
let client = BedrockClient::new(&config);
Ok(Self {
client,
model: model.into(),
region: region_str,
sampling: SamplingParams::default(),
})
}
#[cfg(test)]
pub fn from_client(
client: BedrockClient,
model: impl Into<String>,
region: impl Into<String>,
) -> Self {
Self {
client,
model: model.into(),
region: region.into(),
sampling: SamplingParams::default(),
}
}
pub fn with_sampling(mut self, sampling: SamplingParams) -> Self {
self.sampling = sampling;
self
}
pub fn region(&self) -> &str {
&self.region
}
}
#[async_trait]
impl ChatProvider for BedrockProvider {
fn name(&self) -> &str {
"bedrock"
}
fn model(&self) -> &str {
&self.model
}
async fn chat_stream(
&self,
messages: Vec<ChatMessage>,
_tools: Vec<ToolDef>,
tx: Sender<ChatEvent>,
) -> Result<()> {
let mut system_blocks: Vec<SystemContentBlock> = Vec::new();
let mut converse_messages: Vec<Message> = Vec::new();
for msg in &messages {
if msg.role == "system" {
system_blocks.push(SystemContentBlock::Text(msg.content.clone()));
} else {
let role = if msg.role == "assistant" {
ConversationRole::Assistant
} else {
ConversationRole::User
};
let bedrock_msg = Message::builder()
.role(role)
.content(ContentBlock::Text(msg.content.clone()))
.build()
.context("build Bedrock Message")?;
converse_messages.push(bedrock_msg);
}
}
if converse_messages.is_empty() {
return Err(anyhow!(
"BedrockProvider::chat_stream: no user/assistant messages provided"
));
}
let inference = build_inference_config(&self.sampling);
let mut req = self
.client
.converse_stream()
.model_id(&self.model)
.inference_config(inference)
.set_messages(Some(converse_messages));
if !system_blocks.is_empty() {
req = req.set_system(Some(system_blocks));
}
let output = req.send().await.with_context(|| {
format!(
"AWS Bedrock ConverseStream request failed (model={}, region={}). \
Ensure AWS credentials are configured for Bedrock \
(AWS_ACCESS_KEY_ID / AWS_SECRET_ACCESS_KEY / AWS_PROFILE / IAM role).",
self.model, self.region
)
})?;
let mut stream = output.stream;
loop {
let recv_result = stream
.recv()
.await
.map_err(|sdk_err| format!("Bedrock ConverseStream error: {sdk_err}"));
match handle_stream_event(recv_result, &tx).await {
Flow::Continue => {}
Flow::Stop => return Ok(()),
Flow::Failed(message) => return Err(anyhow!("{message}")),
}
}
}
}
fn build_inference_config(sampling: &SamplingParams) -> InferenceConfiguration {
let stop_sequences = (!sampling.stop.is_empty()).then(|| sampling.stop.clone());
InferenceConfiguration::builder()
.max_tokens(sampling.max_tokens.unwrap_or(4096) as i32)
.set_temperature(sampling.temperature)
.set_stop_sequences(stop_sequences)
.build()
}
#[derive(Debug, PartialEq, Eq)]
enum Flow {
Continue,
Stop,
Failed(String),
}
async fn handle_stream_event(
result: std::result::Result<Option<StreamEvent>, String>,
tx: &Sender<ChatEvent>,
) -> Flow {
match result {
Ok(Some(StreamEvent::ContentBlockDelta(ev))) => {
if let Some(ContentBlockDelta::Text(text)) = ev.delta()
&& tx.send(ChatEvent::Delta(text.clone())).await.is_err()
{
return Flow::Stop;
}
Flow::Continue
}
Ok(Some(StreamEvent::Metadata(meta))) => {
if let Some(usage) = meta.usage() {
let chat_usage = ChatUsage {
prompt_tokens: usage.input_tokens().max(0) as u32,
completion_tokens: usage.output_tokens().max(0) as u32,
cache_read_tokens: usage.cache_read_input_tokens().unwrap_or(0).max(0) as u32,
cache_creation_tokens: usage.cache_write_input_tokens().unwrap_or(0).max(0)
as u32,
};
if tx.send(ChatEvent::Usage(chat_usage)).await.is_err() {
return Flow::Stop;
}
}
Flow::Continue
}
Ok(Some(_)) => Flow::Continue,
Ok(None) => {
let _ = tx.send(ChatEvent::Done).await;
Flow::Stop
}
Err(message) => {
let _ = tx.send(ChatEvent::Error(message.clone())).await;
Flow::Failed(message)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn bedrock_provider_reports_metadata() {
let config = aws_config::defaults(BehaviorVersion::latest())
.region(aws_types::region::Region::new("us-east-1"))
.no_credentials()
.load()
.await;
let client = BedrockClient::new(&config);
let provider = BedrockProvider::from_client(client, DEFAULT_BEDROCK_MODEL, "us-east-1");
assert_eq!(provider.name(), "bedrock");
assert_eq!(provider.model(), DEFAULT_BEDROCK_MODEL);
assert_eq!(provider.region(), "us-east-1");
}
#[test]
fn bedrock_region_resolution() {
assert_eq!(
resolve_bedrock_region(Some("eu-west-1")),
"eu-west-1",
"explicit should win"
);
assert_eq!(
resolve_bedrock_region(Some("")),
DEFAULT_BEDROCK_REGION,
"empty explicit should fall through to default"
);
assert_eq!(
resolve_bedrock_region(None),
DEFAULT_BEDROCK_REGION,
"None should return default"
);
}
#[tokio::test]
async fn bedrock_no_credentials_returns_clear_error() {
let config = aws_config::defaults(BehaviorVersion::latest())
.region(aws_types::region::Region::new("us-east-1"))
.no_credentials()
.load()
.await;
let client = BedrockClient::new(&config);
let provider = BedrockProvider::from_client(client, DEFAULT_BEDROCK_MODEL, "us-east-1");
let (tx, _rx) = tokio::sync::mpsc::channel::<ChatEvent>(8);
let result = provider
.chat_stream(
vec![crate::ChatMessage {
role: "user".into(),
content: "hello".into(),
tool_call_id: None,
tool_calls: None,
}],
vec![],
tx,
)
.await;
let err = result.expect_err("should fail without real credentials");
let msg = format!("{err:#}");
assert!(
msg.to_lowercase().contains("bedrock")
|| msg.to_lowercase().contains("credential")
|| msg.to_lowercase().contains("aws"),
"error message should mention Bedrock/credentials; got: {msg}"
);
}
#[tokio::test]
#[ignore = "requires real AWS credentials with bedrock:InvokeModel permission"]
async fn bedrock_live_converse_stream_smoke_test() {
let provider = BedrockProvider::new(DEFAULT_BEDROCK_MODEL, None)
.await
.expect("BedrockProvider::new failed");
let (tx, mut rx) = tokio::sync::mpsc::channel::<ChatEvent>(8);
let handle = tokio::spawn(async move {
provider
.chat_stream(
vec![
crate::ChatMessage {
role: "system".into(),
content: "You are a concise assistant. Reply in plain text.".into(),
tool_call_id: None,
tool_calls: None,
},
crate::ChatMessage {
role: "user".into(),
content: "Say hello in exactly 3 words.".into(),
tool_call_id: None,
tool_calls: None,
},
],
vec![],
tx,
)
.await
});
let mut text = String::new();
let mut saw_done = false;
let mut usage: Option<ChatUsage> = None;
while let Some(ev) = rx.recv().await {
match ev {
ChatEvent::Delta(s) => text.push_str(&s),
ChatEvent::Done => saw_done = true,
ChatEvent::Error(e) => panic!("stream error: {e}"),
ChatEvent::ToolCall(_) => {}
ChatEvent::Usage(u) => usage = Some(u),
}
}
handle
.await
.expect("task panicked")
.expect("chat_stream failed");
assert!(!text.is_empty(), "expected non-empty response");
assert!(saw_done, "expected ChatEvent::Done");
let usage = usage.expect("expected a ChatEvent::Usage before Done (#3767)");
assert!(
usage.prompt_tokens > 0 || usage.completion_tokens > 0,
"expected non-zero usage: {usage:?}"
);
eprintln!("bedrock_live_converse_stream_smoke_test response: {text:?} usage={usage:?}");
}
#[tokio::test]
async fn bedrock_stream_forwards_text_deltas_in_order() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<ChatEvent>(8);
let flow1 = handle_stream_event(Ok(Some(content_delta("He"))), &tx).await;
assert_eq!(flow1, Flow::Continue);
let flow2 = handle_stream_event(Ok(Some(content_delta("llo"))), &tx).await;
assert_eq!(flow2, Flow::Continue);
let flow3 = handle_stream_event(Ok(None), &tx).await;
assert_eq!(flow3, Flow::Stop);
drop(tx);
let mut collected = Vec::new();
while let Some(ev) = rx.recv().await {
collected.push(ev);
}
assert!(matches!(&collected[0], ChatEvent::Delta(d) if d == "He"));
assert!(matches!(&collected[1], ChatEvent::Delta(d) if d == "llo"));
assert!(matches!(collected[2], ChatEvent::Done));
assert_eq!(collected.len(), 3, "no extra events: {collected:?}");
}
#[tokio::test]
async fn bedrock_stream_ignores_structural_events() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<ChatEvent>(8);
let flow = handle_stream_event(
Ok(Some(StreamEvent::MessageStart(
aws_sdk_bedrockruntime::types::MessageStartEvent::builder()
.role(ConversationRole::Assistant)
.build()
.unwrap(),
))),
&tx,
)
.await;
assert_eq!(flow, Flow::Continue);
let flow = handle_stream_event(
Ok(Some(StreamEvent::MessageStop(
aws_sdk_bedrockruntime::types::MessageStopEvent::builder()
.stop_reason(aws_sdk_bedrockruntime::types::StopReason::EndTurn)
.build()
.unwrap(),
))),
&tx,
)
.await;
assert_eq!(flow, Flow::Continue);
drop(tx);
assert!(
rx.recv().await.is_none(),
"structural events must not emit any ChatEvent"
);
}
#[tokio::test]
async fn bedrock_stream_reports_usage_from_metadata_event() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<ChatEvent>(8);
let flow = handle_stream_event(Ok(Some(metadata_with_usage(120, 45, 30, 10))), &tx).await;
assert_eq!(flow, Flow::Continue);
drop(tx);
let events: Vec<_> = {
let mut out = Vec::new();
while let Some(ev) = rx.recv().await {
out.push(ev);
}
out
};
assert_eq!(events.len(), 1, "expected exactly one event: {events:?}");
match &events[0] {
ChatEvent::Usage(u) => {
assert_eq!(u.prompt_tokens, 120);
assert_eq!(u.completion_tokens, 45);
assert_eq!(u.cache_read_tokens, 30);
assert_eq!(u.cache_creation_tokens, 10);
}
other => panic!("expected ChatEvent::Usage, got {other:?}"),
}
}
#[tokio::test]
async fn bedrock_stream_metadata_without_usage_emits_nothing() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<ChatEvent>(8);
let meta = aws_sdk_bedrockruntime::types::ConverseStreamMetadataEvent::builder().build();
let flow = handle_stream_event(Ok(Some(StreamEvent::Metadata(meta))), &tx).await;
assert_eq!(flow, Flow::Continue);
drop(tx);
assert!(rx.recv().await.is_none(), "no usage means no ChatEvent");
}
#[tokio::test]
async fn bedrock_stream_surfaces_mid_stream_error() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<ChatEvent>(8);
let flow1 = handle_stream_event(Ok(Some(content_delta("partial "))), &tx).await;
assert_eq!(flow1, Flow::Continue);
let flow2 = handle_stream_event(
Err("Bedrock ConverseStream error: throttled".to_string()),
&tx,
)
.await;
assert_eq!(
flow2,
Flow::Failed("Bedrock ConverseStream error: throttled".to_string())
);
drop(tx);
let mut collected = Vec::new();
while let Some(ev) = rx.recv().await {
collected.push(ev);
}
assert!(matches!(&collected[0], ChatEvent::Delta(d) if d == "partial "));
match &collected[1] {
ChatEvent::Error(msg) => assert!(msg.contains("throttled")),
other => panic!("expected ChatEvent::Error, got {other:?}"),
}
assert_eq!(
collected.len(),
2,
"a mid-stream error must not also emit Done: {collected:?}"
);
}
#[tokio::test]
async fn bedrock_stream_done_emits_terminal_marker() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<ChatEvent>(8);
let flow = handle_stream_event(Ok(None), &tx).await;
assert_eq!(flow, Flow::Stop);
drop(tx);
let mut collected = Vec::new();
while let Some(ev) = rx.recv().await {
collected.push(ev);
}
assert_eq!(collected.len(), 1);
assert!(matches!(collected[0], ChatEvent::Done));
}
#[test]
fn bedrock_stream_forwards_sampling_params() {
let sampling = SamplingParams {
temperature: Some(0.2),
max_tokens: Some(512),
stop: vec!["STOP".to_string()],
};
let inference = build_inference_config(&sampling);
assert_eq!(inference.max_tokens(), Some(512));
assert_eq!(inference.temperature(), Some(0.2));
assert_eq!(inference.stop_sequences(), &["STOP".to_string()][..]);
}
#[test]
fn bedrock_stream_sampling_defaults_when_unset() {
let inference = build_inference_config(&SamplingParams::default());
assert_eq!(inference.max_tokens(), Some(4096));
assert_eq!(inference.temperature(), None);
assert!(inference.stop_sequences().is_empty());
}
fn content_delta(text: &str) -> StreamEvent {
StreamEvent::ContentBlockDelta(
aws_sdk_bedrockruntime::types::ContentBlockDeltaEvent::builder()
.delta(ContentBlockDelta::Text(text.to_string()))
.content_block_index(0)
.build()
.expect("build ContentBlockDeltaEvent"),
)
}
fn metadata_with_usage(
input_tokens: i32,
output_tokens: i32,
cache_read: i32,
cache_write: i32,
) -> StreamEvent {
let usage = aws_sdk_bedrockruntime::types::TokenUsage::builder()
.input_tokens(input_tokens)
.output_tokens(output_tokens)
.total_tokens(input_tokens + output_tokens)
.cache_read_input_tokens(cache_read)
.cache_write_input_tokens(cache_write)
.build()
.expect("build TokenUsage");
StreamEvent::Metadata(
aws_sdk_bedrockruntime::types::ConverseStreamMetadataEvent::builder()
.usage(usage)
.build(),
)
}
}