1use super::mantle::{MantleAuth, MantleClient};
2use super::mappers::{default_cache_point, map_messages, map_tools};
3use super::streaming::process_bedrock_stream;
4use crate::catalog::transport::ModelTransport;
5use crate::provider::{LlmResponseStream, ProviderFactory, StreamingModelProvider, get_context_window, stream_from};
6use crate::providers::openai_responses::transport::process_connection;
7use crate::{Context, LlmError, ProviderAuthMode, ProviderConnectionConfig, ProviderError, Result};
8use aws_config::Region;
9use aws_credential_types::provider::SharedCredentialsProvider;
10use aws_sdk_bedrockruntime::config::{BehaviorVersion, Credentials};
11use aws_sdk_bedrockruntime::error::SdkError;
12use aws_sdk_bedrockruntime::operation::converse_stream::ConverseStreamError;
13use aws_sdk_bedrockruntime::primitives::event_stream::EventReceiver;
14use aws_sdk_bedrockruntime::types::error::ConverseStreamOutputError;
15use aws_sdk_bedrockruntime::types::{ConverseStreamOutput, InferenceConfiguration};
16use aws_sdk_bedrockruntime::{Client, Config};
17use tracing::{error, info, warn};
18
19const DEFAULT_MODEL: &str = "anthropic.claude-sonnet-4-5-20250929-v1:0";
20const DEFAULT_MAX_TOKENS: i32 = 16_384;
21const DEFAULT_REGION: &str = "us-east-1";
22
23#[derive(Clone)]
25pub struct AwsCredentials {
26 pub access_key_id: String,
27 pub secret_access_key: String,
28 pub session_token: Option<String>,
29}
30
31#[derive(Clone)]
32pub struct BedrockProvider {
33 client: Client,
34 mantle: MantleClient,
35 model: String,
36 inference_profile_arn: Option<String>,
37}
38
39impl BedrockProvider {
40 pub async fn new(connection: ProviderConnectionConfig) -> Self {
43 if connection.auth_mode == ProviderAuthMode::None {
44 return Self::from_config(None, region_from_env().as_deref(), connection);
45 }
46
47 let mut loader = aws_config::defaults(BehaviorVersion::latest());
48 if let Some(url) = &connection.base_url {
49 loader = loader.endpoint_url(url.clone());
50 }
51
52 let config = loader.load().await;
53 let region = config
54 .region()
55 .map(ToString::to_string)
56 .or_else(region_from_env)
57 .unwrap_or_else(|| DEFAULT_REGION.to_string());
58 let auth = mantle_auth(config.credentials_provider(), ®ion);
59 Self::assemble(Client::new(&config), region, auth, connection)
60 }
61
62 pub fn from_config(
64 credentials: Option<AwsCredentials>,
65 region: Option<&str>,
66 connection: ProviderConnectionConfig,
67 ) -> Self {
68 let region = region.unwrap_or(DEFAULT_REGION).to_string();
69 let auth = match connection.auth_mode {
70 ProviderAuthMode::Default => mantle_auth(credentials.clone().map(shared_credentials), ®ion),
71 ProviderAuthMode::None => MantleAuth::None,
72 };
73 let client = build_client(credentials, ®ion, connection.base_url.as_deref(), connection.auth_mode);
74 Self::assemble(client, region, auth, connection)
75 }
76
77 fn assemble(client: Client, region: String, auth: MantleAuth, connection: ProviderConnectionConfig) -> Self {
78 let mantle = MantleClient::new(region, auth, connection.base_url.clone());
79 Self {
80 client,
81 mantle,
82 model: DEFAULT_MODEL.to_string(),
83 inference_profile_arn: connection.inference_profile_arn,
84 }
85 }
86
87 pub fn with_model(mut self, model: &str) -> Self {
88 self.model = model.to_string();
89 self
90 }
91
92 pub fn with_inference_profile_arn(mut self, arn: impl Into<String>) -> Self {
93 self.inference_profile_arn = Some(arn.into());
94 self
95 }
96
97 pub fn with_bearer_token(mut self, token: impl Into<String>) -> Self {
100 self.mantle = self.mantle.with_auth(MantleAuth::BearerToken(token.into()));
101 self
102 }
103
104 fn request_model_id(&self) -> &str {
105 self.inference_profile_arn.as_deref().unwrap_or(&self.model)
106 }
107
108 fn mantle_transport(&self) -> Option<ModelTransport> {
111 self.model().and_then(|model| model.transport())
112 }
113
114 async fn send_converse_stream(
115 &self,
116 context: &Context,
117 ) -> Result<EventReceiver<ConverseStreamOutput, ConverseStreamOutputError>> {
118 let cache_point =
119 self.model().is_some_and(|m| m.supports_prompt_caching()).then(default_cache_point).transpose()?;
120 let (system_blocks, messages) = map_messages(context.messages(), cache_point.as_ref())?;
121 let settings = context.model_settings();
122 let max_tokens = settings.max_tokens.and_then(|m| i32::try_from(m).ok()).unwrap_or(DEFAULT_MAX_TOKENS);
123 let mut inference_config = InferenceConfiguration::builder().max_tokens(max_tokens);
124
125 if let Some(temp) = settings.temperature {
126 inference_config = inference_config.temperature(temp);
127 }
128
129 if let Some(top_p) = settings.top_p {
130 inference_config = inference_config.top_p(top_p);
131 }
132
133 let inference_config = inference_config.build();
134
135 let mut request = self
136 .client
137 .converse_stream()
138 .model_id(self.request_model_id())
139 .set_messages(Some(messages))
140 .inference_config(inference_config);
141
142 if !system_blocks.is_empty() {
143 request = request.set_system(Some(system_blocks));
144 }
145
146 if !context.tools().is_empty() {
147 let tool_config = map_tools(context.tools(), cache_point.as_ref())?;
148 request = request.tool_config(tool_config);
149 }
150
151 if let Some(arn) = self.inference_profile_arn.as_deref() {
152 info!(model = %self.model, inference_profile_arn = %arn, "Sending Bedrock converse_stream request");
153 } else {
154 info!(model = %self.model, "Sending Bedrock converse_stream request");
155 }
156
157 let response = request.send().await.map_err(|e| {
158 error!(model = %self.model, error = ?e, "Bedrock API error");
159 LlmError::from(e)
160 })?;
161
162 Ok(response.stream)
163 }
164}
165
166impl ProviderFactory for BedrockProvider {
167 async fn from_env() -> Result<Self> {
168 Ok(Self::new(ProviderConnectionConfig::default()).await)
169 }
170
171 async fn from_env_with_connection(connection: ProviderConnectionConfig) -> Result<Self> {
172 Ok(Self::new(connection).await)
173 }
174
175 fn with_model(self, model: &str) -> Self {
176 self.with_model(model)
177 }
178}
179
180impl StreamingModelProvider for BedrockProvider {
181 fn model(&self) -> Option<crate::LlmModel> {
182 format!("bedrock:{}", self.model).parse().ok()
183 }
184
185 fn context_window(&self) -> Option<u32> {
186 get_context_window("bedrock", &self.model)
187 }
188
189 fn stream_response(&self, context: &Context) -> LlmResponseStream {
190 if let Err(error) = crate::provider::validate_reasoning(context, self.model().as_ref()) {
191 return crate::provider::error_stream(error);
192 }
193 let provider = self.clone();
194 let context = context.clone();
195
196 let Some(transport) = self.mantle_transport() else {
197 return stream_from(async move { provider.send_converse_stream(&context).await }, process_bedrock_stream);
198 };
199
200 if let Some(arn) = self.inference_profile_arn.as_deref() {
201 warn!(
202 model = %self.model,
203 inference_profile_arn = %arn,
204 "Ignoring inferenceProfileArn: this model is served by the Responses API, which has no inference profiles"
205 );
206 }
207
208 stream_from(
209 async move { provider.mantle.stream(&provider.model, &transport, &context).await },
210 process_connection,
211 )
212 }
213
214 fn display_name(&self) -> String {
215 format!("Bedrock ({})", self.model)
216 }
217}
218
219impl From<SdkError<ConverseStreamError>> for LlmError {
220 fn from(e: SdkError<ConverseStreamError>) -> Self {
221 let message = format!("Bedrock API error: {e}");
222 let status = e.raw_response().map(|r| r.status().as_u16());
223 let request_id = e.raw_response().and_then(|r| r.headers().get("x-amzn-requestid")).map(str::to_string);
224 let mut provider = match &e {
225 SdkError::TimeoutError(_) => ProviderError::timeout(message),
226 SdkError::DispatchFailure(_) => ProviderError::network(message),
227 SdkError::ResponseError(_) => ProviderError::server(message),
228 SdkError::ServiceError(svc) => {
229 let inner = svc.err();
230 if inner.is_throttling_exception() {
231 ProviderError::rate_limit(message)
232 } else if inner.is_service_unavailable_exception()
233 || inner.is_internal_server_exception()
234 || inner.is_model_stream_error_exception()
235 {
236 ProviderError::server(message)
237 } else {
238 ProviderError::api(message)
239 }
240 }
241 _ => ProviderError::api(message),
242 };
243 provider = provider.with_http_metadata(status, request_id);
244 Self::from(provider)
245 }
246}
247
248fn build_client(
249 credentials: Option<AwsCredentials>,
250 region: &str,
251 base_url: Option<&str>,
252 auth_mode: ProviderAuthMode,
253) -> Client {
254 let mut config =
255 Config::builder().behavior_version(BehaviorVersion::latest()).region(Region::new(region.to_string()));
256
257 if auth_mode == ProviderAuthMode::None {
258 config = config.allow_no_auth();
259 } else if let Some(credentials) = credentials {
260 config = config.credentials_provider(shared_credentials(credentials));
261 }
262 if let Some(url) = base_url {
263 config = config.endpoint_url(url);
264 }
265
266 Client::from_conf(config.build())
267}
268
269fn shared_credentials(credentials: AwsCredentials) -> SharedCredentialsProvider {
270 SharedCredentialsProvider::new(Credentials::new(
271 credentials.access_key_id,
272 credentials.secret_access_key,
273 credentials.session_token,
274 None,
275 "aether-bedrock-provider",
276 ))
277}
278
279fn mantle_auth(credentials: Option<SharedCredentialsProvider>, region: &str) -> MantleAuth {
285 if let Some(token) = MantleAuth::bearer_token_from_env() {
286 return MantleAuth::BearerToken(token);
287 }
288 match credentials {
289 Some(credentials) => MantleAuth::SigV4 { credentials, region: region.to_string() },
290 None => MantleAuth::None,
291 }
292}
293
294fn region_from_env() -> Option<String> {
295 ["AWS_REGION", "AWS_DEFAULT_REGION"].into_iter().find_map(|name| match std::env::var(name) {
296 Ok(value) if !value.is_empty() => Some(value),
297 _ => None,
298 })
299}
300
301#[cfg(test)]
302mod tests {
303 use super::*;
304 use crate::catalog::Provider;
305 use crate::providers::test_capture_server::CaptureServer;
306 use crate::types::IsoString;
307 use crate::{AssistantReasoning, ChatMessage, EncryptedReasoningContent, LlmModel, MessageId};
308 use axum::Router;
309 use axum::body::Body;
310 use axum::extract::State;
311 use axum::http::{HeaderMap, Method, Request, StatusCode};
312 use axum::response::IntoResponse;
313 use axum::routing::any;
314 use futures::StreamExt;
315 use std::sync::Arc;
316 use tokio::net::TcpListener;
317 use tokio::sync::{Mutex, oneshot};
318 use utils::ReasoningEffort;
319
320 fn inference_profile_arn(model: &str) -> String {
321 format!("arn:aws:bedrock:us-west-2:000000000000:inference-profile/{model}")
322 }
323
324 fn application_inference_profile_arn() -> &'static str {
325 "arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000"
326 }
327
328 fn test_provider() -> BedrockProvider {
329 BedrockProvider::from_config(None, None, ProviderConnectionConfig::default())
330 }
331
332 fn mantle_model() -> String {
336 LlmModel::all()
337 .iter()
338 .find(|model| model.provider_enum() == Provider::Bedrock && model.transport().is_some())
339 .expect("catalog must expose at least one Responses-transport Bedrock model")
340 .model_id()
341 .to_string()
342 }
343
344 async fn mantle_provider(server: &CaptureServer) -> BedrockProvider {
346 BedrockProvider::new(ProviderConnectionConfig {
347 base_url: Some(server.base_url.clone()),
348 auth_mode: ProviderAuthMode::None,
349 ..Default::default()
350 })
351 .await
352 .with_model(&mantle_model())
353 }
354
355 #[test]
356 fn test_display_name() {
357 assert_eq!(test_provider().display_name(), "Bedrock (anthropic.claude-sonnet-4-5-20250929-v1:0)");
358 }
359
360 #[test]
361 fn test_with_model() {
362 let provider = test_provider().with_model("anthropic.claude-opus-4-20250514-v1:0");
363 assert_eq!(provider.display_name(), "Bedrock (anthropic.claude-opus-4-20250514-v1:0)");
364 }
365
366 #[test]
367 fn test_default_values() {
368 let provider = test_provider();
369 assert_eq!(provider.model, "anthropic.claude-sonnet-4-5-20250929-v1:0");
370 }
371
372 #[tokio::test]
373 async fn auth_none_sends_unsigned_request_to_custom_endpoint() {
374 let endpoint = FakeBedrockEndpoint::start().await;
375 let provider = BedrockProvider::new(ProviderConnectionConfig {
376 base_url: Some(endpoint.url.clone()),
377 auth_mode: ProviderAuthMode::None,
378 ..Default::default()
379 })
380 .await;
381
382 let result = provider.send_converse_stream(&hello_context()).await;
383 let request = endpoint.request.await.expect("fake Bedrock endpoint received no request");
384
385 assert!(result.is_err());
386 assert_eq!(request.method, Method::POST);
387 assert!(request.path.starts_with("/model/"), "{}", request.path);
388 assert!(!request.headers.contains_key("authorization"), "request was signed: {:?}", request.headers);
389 assert!(
390 !request.headers.contains_key("x-amz-security-token"),
391 "request included session token: {:?}",
392 request.headers
393 );
394 }
395
396 fn static_credentials() -> AwsCredentials {
397 AwsCredentials {
398 access_key_id: "AKIAIOSFODNN7EXAMPLE".to_string(),
399 secret_access_key: "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY".to_string(),
400 session_token: None,
401 }
402 }
403
404 fn hello_context() -> Context {
405 Context::new(vec![ChatMessage::user("Hello")], vec![])
406 }
407
408 #[tokio::test]
409 async fn mantle_disabled_uses_none_without_summary() {
410 let model = LlmModel::all()
411 .iter()
412 .find(|model| {
413 model.provider_enum() == Provider::Bedrock
414 && model.transport().is_some()
415 && model.supports_reasoning_off()
416 })
417 .unwrap();
418 let mut server = CaptureServer::start_responses().await;
419 let provider = mantle_provider(&server).await.with_model(&model.model_id());
420 let mut context = hello_context();
421 context.set_reasoning_effort(ReasoningEffort::Disabled);
422 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
423 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
424 let body = server.captured().await.body;
425 assert_eq!(body["reasoning"]["effort"], "none");
426 assert!(body["reasoning"]["summary"].is_null());
427 }
428
429 #[tokio::test]
430 async fn converse_disabled_is_a_non_retryable_transport_error() {
431 let model = LlmModel::all()
432 .iter()
433 .find(|model| {
434 model.provider_enum() == Provider::Bedrock
435 && model.transport().is_none()
436 && model.supports_reasoning_off()
437 })
438 .unwrap();
439 let provider =
440 BedrockProvider::new(ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() })
441 .await
442 .with_model(&model.model_id());
443 let mut context = hello_context();
444 context.set_reasoning_effort(crate::ReasoningEffort::Disabled);
445 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
446 assert_eq!(responses.len(), 1);
447 assert!(matches!(&responses[0], Err(LlmError::UnsupportedDisableTransport { .. })));
448 assert!(!responses[0].as_ref().unwrap_err().is_retryable());
449 }
450
451 #[tokio::test]
452 async fn responses_shape_models_are_sent_to_the_responses_endpoint() {
453 let mut server = CaptureServer::start_responses().await;
454 let provider = mantle_provider(&server).await;
455 let mut context = hello_context();
456 context.set_reasoning_effort(crate::ReasoningEffort::Xhigh);
457
458 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
459 let captured = server.captured().await;
460
461 assert!(!responses.is_empty());
462 assert_eq!(captured.path, "/responses");
463 assert_eq!(captured.body["model"], mantle_model());
464 assert_eq!(captured.body["stream"], true);
465 assert_eq!(captured.body["store"], false);
466 assert_eq!(captured.body["reasoning"]["effort"], "xhigh");
467 assert_eq!(captured.body["input"][0]["role"], "user");
468 assert_eq!(captured.body["input"][0]["content"][0]["type"], "input_text");
469 assert_eq!(captured.body["input"][0]["content"][0]["text"], "Hello");
470 assert!(captured.headers.get("authorization").is_none(), "{:?}", captured.headers);
471 }
472
473 #[tokio::test]
474 async fn http_200_failed_server_error_is_retryable_with_diagnostics() {
475 use crate::providers::test_capture_server::ResponseSpec;
476 let spec = ResponseSpec::sse(include_str!("../../../tests/fixtures/openai_responses/04_failed_server.sse"))
477 .with_header("x-amzn-requestid", "amzn-req-123");
478 let mut server = CaptureServer::start_with_spec(spec).await;
479 let provider = mantle_provider(&server).await;
480
481 let responses = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
482 let _ = server.captured().await;
483
484 assert!(!responses.iter().any(|r| matches!(r, Ok(crate::LlmResponse::Done { .. }))));
485 let err = responses.iter().find_map(|r| r.as_ref().err()).expect("expected a failure");
486 assert!(err.is_retryable(), "server_error must be retryable: {err:?}");
487 let provider_error = err.provider().expect("expected provider error");
488 assert_eq!(provider_error.kind, crate::ProviderErrorKind::Server);
489 assert_eq!(provider_error.code.as_deref(), Some("server_error"));
490 assert_eq!(provider_error.http_status, Some(200));
491 assert_eq!(provider_error.request_id.as_deref(), Some("amzn-req-123"));
492 }
493
494 #[tokio::test]
495 async fn failed_event_with_unknown_code_is_terminal_without_done() {
496 let body = "event: response.created\ndata: {\"type\":\"response.created\",\"response\":{\"id\":\"r1\"}}\n\nevent: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"code\":\"invalid_prompt\",\"message\":\"bad prompt\"}}}\n\n";
497 let mut server = CaptureServer::start_with_response(body).await;
498 let provider = mantle_provider(&server).await;
499
500 let responses = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
501 let _ = server.captured().await;
502
503 assert!(!responses.iter().any(|r| matches!(r, Ok(crate::LlmResponse::Done { .. }))));
504 let err = responses.iter().find_map(|r| r.as_ref().err()).expect("expected a failure");
505 assert!(!err.is_retryable(), "invalid_prompt must be terminal: {err:?}");
506 }
507
508 #[tokio::test]
509 async fn responses_transport_rejects_malformed_and_truncated_streams() {
510 for response in [
511 "data: {not-json}\n\n",
512 "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_1\"}}\n\ndata: [DONE]\n\n",
513 ] {
514 let mut server = CaptureServer::start_with_response(response).await;
515 let provider = mantle_provider(&server).await;
516
517 let responses = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
518 let _ = server.captured().await;
519
520 assert!(responses.iter().any(|response| {
521 response.as_ref().err().and_then(LlmError::provider).map(|provider| provider.kind)
522 == Some(crate::ProviderErrorKind::StreamInterrupted)
523 }));
524 assert!(!responses.iter().any(|response| matches!(response, Ok(crate::LlmResponse::Done { .. }))));
525 }
526 }
527
528 #[tokio::test]
529 async fn responses_transport_drops_encrypted_reasoning_from_another_model() {
530 let mut server = CaptureServer::start_responses().await;
531 let provider = mantle_provider(&server).await;
532 let context = Context::new(
533 vec![ChatMessage::Assistant {
534 message_id: MessageId::new(),
535 content: "previous answer".to_string(),
536 reasoning: AssistantReasoning {
537 summary_text: None,
538 encrypted_content: Some(EncryptedReasoningContent {
539 id: "reasoning-id".to_string(),
540 model: "bedrock:openai.gpt-5.5".parse().unwrap(),
541 content: "opaque-for-another-model".to_string(),
542 }),
543 },
544 timestamp: IsoString::now(),
545 tool_calls: vec![],
546 }],
547 vec![],
548 );
549
550 let _ = provider.stream_response(&context).collect::<Vec<_>>().await;
551 let captured = server.captured().await;
552
553 assert!(
554 captured.body["input"].as_array().unwrap().iter().all(|item| item["type"] != "reasoning"),
555 "{}",
556 captured.body
557 );
558 }
559
560 #[tokio::test]
561 async fn converse_shape_models_do_not_use_the_responses_endpoint() {
562 let endpoint = FakeBedrockEndpoint::start().await;
563 let provider = BedrockProvider::new(ProviderConnectionConfig {
564 base_url: Some(endpoint.url.clone()),
565 auth_mode: ProviderAuthMode::None,
566 ..Default::default()
567 })
568 .await
569 .with_model(DEFAULT_MODEL);
570
571 let _ = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
572 let request = endpoint.request.await.expect("fake Bedrock endpoint received no request");
573
574 assert!(request.path.starts_with("/model/"), "{}", request.path);
575 }
576
577 #[tokio::test]
578 async fn bearer_token_authenticates_responses_requests() {
579 let mut server = CaptureServer::start_responses().await;
580 let provider = BedrockProvider::from_config(
581 None,
582 Some("us-west-2"),
583 ProviderConnectionConfig { base_url: Some(server.base_url.clone()), ..Default::default() },
584 )
585 .with_bearer_token("test-token")
586 .with_model(&mantle_model());
587
588 let _ = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
589 let captured = server.captured().await;
590
591 assert_eq!(captured.headers.get("authorization").unwrap(), "Bearer test-token");
592 }
593
594 #[tokio::test]
595 async fn credential_chain_sigv4_signs_responses_requests() {
596 let mut server = CaptureServer::start_responses().await;
597 let provider = BedrockProvider::from_config(
598 Some(static_credentials()),
599 Some("us-west-2"),
600 ProviderConnectionConfig { base_url: Some(server.base_url.clone()), ..Default::default() },
601 )
602 .with_model(&mantle_model());
603
604 let _ = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
605 let captured = server.captured().await;
606
607 let authorization = captured.headers.get("authorization").expect("request was not signed").to_str().unwrap();
608 assert!(
609 authorization.starts_with("AWS4-HMAC-SHA256 Credential=AKIAIOSFODNN7EXAMPLE/"),
610 "unexpected authorization header: {authorization}"
611 );
612 assert!(authorization.contains("/us-west-2/bedrock/aws4_request"), "{authorization}");
613 assert!(captured.headers.contains_key("x-amz-date"), "{:?}", captured.headers);
614 }
615
616 #[test]
617 fn only_responses_shape_models_route_to_the_mantle_transport() {
618 let mantle = test_provider().with_model(&mantle_model());
619 let converse = test_provider().with_model(DEFAULT_MODEL);
620 let profile = test_provider().with_model("us.anthropic.claude-future-model-v99:0");
621
622 assert!(matches!(mantle.mantle_transport(), Some(ModelTransport::OpenAiResponses { .. })));
623 assert_eq!(converse.mantle_transport(), None);
624 assert_eq!(profile.mantle_transport(), None);
625 }
626
627 #[test]
628 fn explicit_connection_preserves_inference_profile() {
629 let provider = BedrockProvider::from_config(
630 None,
631 Some("us-west-2"),
632 ProviderConnectionConfig { inference_profile_arn: Some("arn:test".to_string()), ..Default::default() },
633 );
634
635 assert_eq!(provider.inference_profile_arn.as_deref(), Some("arn:test"));
636 }
637
638 #[test]
639 fn test_from_config_with_credentials() {
640 let provider =
641 BedrockProvider::from_config(Some(static_credentials()), None, ProviderConnectionConfig::default());
642 assert_eq!(provider.model, DEFAULT_MODEL);
643 }
644
645 #[test]
646 fn test_from_config_with_credentials_and_region() {
647 let credentials =
648 AwsCredentials { session_token: Some("FwoGZXIvYXdzEBYaD...".to_string()), ..static_credentials() };
649
650 let provider =
651 BedrockProvider::from_config(Some(credentials), Some("us-west-2"), ProviderConnectionConfig::default())
652 .with_model("anthropic.claude-opus-4-20250514-v1:0");
653
654 assert_eq!(provider.model, "anthropic.claude-opus-4-20250514-v1:0");
655 }
656
657 #[test]
658 fn test_from_config_with_region_only() {
659 let provider = BedrockProvider::from_config(None, Some("eu-west-1"), ProviderConnectionConfig::default());
660 assert_eq!(provider.model, DEFAULT_MODEL);
661 }
662
663 #[test]
664 fn catalog_foundation_id_resolves_context_window() {
665 let provider = test_provider().with_model("anthropic.claude-sonnet-4-5-20250929-v1:0");
666 assert!(provider.context_window().is_some());
667 assert_eq!(provider.model().unwrap().to_string(), "bedrock:anthropic.claude-sonnet-4-5-20250929-v1:0");
668 }
669
670 #[test]
671 fn cross_region_profile_id_in_catalog_resolves() {
672 let provider = test_provider().with_model("us.anthropic.claude-opus-4-6-v1");
673 assert!(provider.context_window().is_some());
674 }
675
676 #[test]
677 fn unknown_cross_region_profile_id_falls_through_to_profile() {
678 let id = "us.anthropic.claude-future-model-v99:0";
679 let provider = test_provider().with_model(id);
680 assert_eq!(provider.context_window(), None);
681 assert_eq!(provider.model().unwrap().to_string(), format!("bedrock:{id}"));
682 assert_eq!(provider.display_name(), format!("Bedrock ({id})"));
683 }
684
685 #[tokio::test]
686 async fn separate_inference_profile_arn_is_used_as_request_model_id() {
687 let endpoint = FakeBedrockEndpoint::start().await;
688 let provider = BedrockProvider::new(ProviderConnectionConfig {
689 base_url: Some(endpoint.url.clone()),
690 auth_mode: ProviderAuthMode::None,
691 request_model: None,
692 inference_profile_arn: Some(application_inference_profile_arn().to_string()),
693 })
694 .await
695 .with_model(DEFAULT_MODEL);
696
697 let result = provider.send_converse_stream(&hello_context()).await;
698 let request = endpoint.request.await.expect("fake Bedrock endpoint received no request");
699
700 assert!(result.is_err());
701 assert!(
702 request.path.contains("arn%3Aaws%3Abedrock%3Aus-west-2%3A000000000000%3Aapplication-inference-profile"),
703 "{}",
704 request.path
705 );
706 assert_eq!(provider.context_window(), Some(200_000));
707 assert_eq!(provider.model().unwrap().to_string(), "bedrock:anthropic.claude-sonnet-4-5-20250929-v1:0");
708 }
709
710 #[test]
711 fn with_inference_profile_arn_keeps_canonical_model_identity() {
712 let arn = inference_profile_arn("us.anthropic.claude-sonnet-4-5-20250929-v1:0");
713 let provider =
714 test_provider().with_model("anthropic.claude-sonnet-4-5-20250929-v1:0").with_inference_profile_arn(&arn);
715
716 assert_eq!(provider.request_model_id(), arn);
717 assert_eq!(provider.context_window(), Some(200_000));
718 assert_eq!(provider.model().unwrap().to_string(), "bedrock:anthropic.claude-sonnet-4-5-20250929-v1:0");
719 }
720
721 #[test]
722 fn prompt_caching_support_comes_from_canonical_model() {
723 let cached = test_provider().with_model("anthropic.claude-sonnet-4-5-20250929-v1:0");
724 assert!(cached.model().unwrap().supports_prompt_caching());
725
726 let unknown_profile = test_provider().with_model("us.anthropic.claude-future-model-v99:0");
727 assert!(!unknown_profile.model().unwrap().supports_prompt_caching());
728 }
729
730 struct FakeBedrockEndpoint {
731 url: String,
732 request: oneshot::Receiver<CapturedRequest>,
733 }
734
735 struct CapturedRequest {
736 method: Method,
737 path: String,
738 headers: HeaderMap,
739 }
740
741 #[derive(Clone)]
742 struct FakeBedrockState {
743 request_tx: Arc<Mutex<Option<oneshot::Sender<CapturedRequest>>>>,
744 shutdown_tx: Arc<Mutex<Option<oneshot::Sender<()>>>>,
745 }
746
747 impl FakeBedrockEndpoint {
748 async fn start() -> Self {
749 let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind fake Bedrock endpoint");
750 let url = format!("http://{}", listener.local_addr().expect("fake Bedrock endpoint address"));
751 let (request_tx, request) = oneshot::channel();
752 let (shutdown_tx, shutdown) = oneshot::channel();
753 let state = FakeBedrockState {
754 request_tx: Arc::new(Mutex::new(Some(request_tx))),
755 shutdown_tx: Arc::new(Mutex::new(Some(shutdown_tx))),
756 };
757
758 let app = Router::new().fallback(any(capture_bedrock_request)).with_state(state);
759 tokio::spawn(async move {
760 axum::serve(listener, app)
761 .with_graceful_shutdown(async {
762 let _ = shutdown.await;
763 })
764 .await
765 .expect("serve fake Bedrock endpoint");
766 });
767
768 Self { url, request }
769 }
770 }
771
772 async fn capture_bedrock_request(
773 State(state): State<FakeBedrockState>,
774 request: Request<Body>,
775 ) -> impl IntoResponse {
776 let (parts, _) = request.into_parts();
777 if let Some(tx) = state.request_tx.lock().await.take() {
778 let _ = tx.send(CapturedRequest {
779 method: parts.method,
780 path: parts.uri.path().to_string(),
781 headers: parts.headers,
782 });
783 }
784 if let Some(tx) = state.shutdown_tx.lock().await.take() {
785 let _ = tx.send(());
786 }
787 (StatusCode::FORBIDDEN, "{}")
788 }
789}