Skip to main content

llm/providers/bedrock/
provider.rs

1use super::mantle::{MantleAuth, MantleClient};
2use super::mappers::{default_cache_point, map_messages, map_tools};
3use super::streaming::{converse_events, 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/// AWS credentials for explicit authentication with Bedrock.
24#[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    /// Create a provider using the default AWS credential chain
41    /// (env vars, `~/.aws/credentials`, IAM roles, SSO).
42    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(), &region);
59        Self::assemble(Client::new(&config), region, auth, connection)
60    }
61
62    /// Create a provider from explicit configuration without async credential discovery.
63    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), &region),
71            ProviderAuthMode::None => MantleAuth::None,
72        };
73        let client = build_client(credentials, &region, 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    /// Authenticate Responses-transport requests with a Bedrock API key instead
98    /// of `SigV4`. Equivalent to setting `AWS_BEARER_TOKEN_BEDROCK`.
99    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    /// The Responses-API transport for the current model, when the catalog says
109    /// it is not served by the Converse API.
110    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 }, |receiver| {
198                process_bedrock_stream(converse_events(receiver))
199            });
200        };
201
202        if let Some(arn) = self.inference_profile_arn.as_deref() {
203            warn!(
204                model = %self.model,
205                inference_profile_arn = %arn,
206                "Ignoring inferenceProfileArn: Responses requests use the selected model ID"
207            );
208        }
209
210        stream_from(
211            async move { provider.mantle.stream(&provider.model, &transport, &context).await },
212            process_connection,
213        )
214    }
215
216    fn display_name(&self) -> String {
217        format!("Bedrock ({})", self.model)
218    }
219}
220
221impl From<SdkError<ConverseStreamError>> for LlmError {
222    fn from(e: SdkError<ConverseStreamError>) -> Self {
223        let message = format!("Bedrock API error: {e}");
224        let status = e.raw_response().map(|r| r.status().as_u16());
225        let request_id = e.raw_response().and_then(|r| r.headers().get("x-amzn-requestid")).map(str::to_string);
226        let mut provider = match &e {
227            SdkError::TimeoutError(_) => ProviderError::timeout(message),
228            SdkError::DispatchFailure(_) => ProviderError::network(message),
229            SdkError::ResponseError(_) => ProviderError::server(message),
230            SdkError::ServiceError(svc) => {
231                let inner = svc.err();
232                if inner.is_throttling_exception() {
233                    ProviderError::rate_limit(message)
234                } else if inner.is_service_unavailable_exception()
235                    || inner.is_internal_server_exception()
236                    || inner.is_model_stream_error_exception()
237                {
238                    ProviderError::server(message)
239                } else {
240                    ProviderError::api(message)
241                }
242            }
243            _ => ProviderError::api(message),
244        };
245        provider = provider.with_http_metadata(status, request_id);
246        Self::from(provider)
247    }
248}
249
250fn build_client(
251    credentials: Option<AwsCredentials>,
252    region: &str,
253    base_url: Option<&str>,
254    auth_mode: ProviderAuthMode,
255) -> Client {
256    let mut config =
257        Config::builder().behavior_version(BehaviorVersion::latest()).region(Region::new(region.to_string()));
258
259    if auth_mode == ProviderAuthMode::None {
260        config = config.allow_no_auth();
261    } else if let Some(credentials) = credentials {
262        config = config.credentials_provider(shared_credentials(credentials));
263    }
264    if let Some(url) = base_url {
265        config = config.endpoint_url(url);
266    }
267
268    Client::from_conf(config.build())
269}
270
271fn shared_credentials(credentials: AwsCredentials) -> SharedCredentialsProvider {
272    SharedCredentialsProvider::new(Credentials::new(
273        credentials.access_key_id,
274        credentials.secret_access_key,
275        credentials.session_token,
276        None,
277        "aether-bedrock-provider",
278    ))
279}
280
281/// Resolve the credential scheme for the Responses transport.
282///
283/// A Bedrock API key takes precedence over the credential chain because it is
284/// the scheme the model catalog advertises for these endpoints; `SigV4` keeps
285/// SSO and IAM-role users working without extra configuration.
286fn mantle_auth(credentials: Option<SharedCredentialsProvider>, region: &str) -> MantleAuth {
287    if let Some(token) = MantleAuth::bearer_token_from_env() {
288        return MantleAuth::BearerToken(token);
289    }
290    match credentials {
291        Some(credentials) => MantleAuth::SigV4 { credentials, region: region.to_string() },
292        None => MantleAuth::None,
293    }
294}
295
296fn region_from_env() -> Option<String> {
297    ["AWS_REGION", "AWS_DEFAULT_REGION"].into_iter().find_map(|name| match std::env::var(name) {
298        Ok(value) if !value.is_empty() => Some(value),
299        _ => None,
300    })
301}
302
303#[cfg(test)]
304mod tests {
305    use super::*;
306    use crate::catalog::Provider;
307    use crate::providers::test_capture_server::CaptureServer;
308    use crate::types::IsoString;
309    use crate::{AssistantReasoning, ChatMessage, EncryptedReasoningContent, LlmModel, MessageId};
310    use axum::Router;
311    use axum::body::Body;
312    use axum::extract::State;
313    use axum::http::{HeaderMap, Method, Request, StatusCode};
314    use axum::response::IntoResponse;
315    use axum::routing::any;
316    use futures::StreamExt;
317    use std::sync::Arc;
318    use tokio::net::TcpListener;
319    use tokio::sync::{Mutex, oneshot};
320    use utils::ReasoningEffort;
321
322    fn inference_profile_arn(model: &str) -> String {
323        format!("arn:aws:bedrock:us-west-2:000000000000:inference-profile/{model}")
324    }
325
326    fn application_inference_profile_arn() -> &'static str {
327        "arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000"
328    }
329
330    fn test_provider() -> BedrockProvider {
331        BedrockProvider::from_config(None, None, ProviderConnectionConfig::default())
332    }
333
334    /// A catalog model routed to the Responses transport, resolved from the
335    /// catalog so a models.dev sync that retires one model does not quietly
336    /// leave these tests exercising the Converse path instead.
337    fn mantle_model() -> String {
338        LlmModel::all()
339            .iter()
340            .find(|model| model.provider_enum() == Provider::Bedrock && model.transport().is_some())
341            .expect("catalog must expose at least one Responses-transport Bedrock model")
342            .model_id()
343            .to_string()
344    }
345
346    /// A provider talking to `server` over the Responses transport, unauthenticated.
347    async fn mantle_provider(server: &CaptureServer) -> BedrockProvider {
348        BedrockProvider::new(ProviderConnectionConfig {
349            base_url: Some(server.base_url.clone()),
350            auth_mode: ProviderAuthMode::None,
351            ..Default::default()
352        })
353        .await
354        .with_model(&mantle_model())
355    }
356
357    #[test]
358    fn test_display_name() {
359        assert_eq!(test_provider().display_name(), "Bedrock (anthropic.claude-sonnet-4-5-20250929-v1:0)");
360    }
361
362    #[test]
363    fn test_with_model() {
364        let provider = test_provider().with_model("anthropic.claude-opus-4-20250514-v1:0");
365        assert_eq!(provider.display_name(), "Bedrock (anthropic.claude-opus-4-20250514-v1:0)");
366    }
367
368    #[test]
369    fn test_default_values() {
370        let provider = test_provider();
371        assert_eq!(provider.model, "anthropic.claude-sonnet-4-5-20250929-v1:0");
372    }
373
374    #[tokio::test]
375    async fn auth_none_sends_unsigned_request_to_custom_endpoint() {
376        let endpoint = FakeBedrockEndpoint::start().await;
377        let provider = BedrockProvider::new(ProviderConnectionConfig {
378            base_url: Some(endpoint.url.clone()),
379            auth_mode: ProviderAuthMode::None,
380            ..Default::default()
381        })
382        .await;
383
384        let result = provider.send_converse_stream(&hello_context()).await;
385        let request = endpoint.request.await.expect("fake Bedrock endpoint received no request");
386
387        assert!(result.is_err());
388        assert_eq!(request.method, Method::POST);
389        assert!(request.path.starts_with("/model/"), "{}", request.path);
390        assert!(!request.headers.contains_key("authorization"), "request was signed: {:?}", request.headers);
391        assert!(
392            !request.headers.contains_key("x-amz-security-token"),
393            "request included session token: {:?}",
394            request.headers
395        );
396    }
397
398    fn static_credentials() -> AwsCredentials {
399        AwsCredentials {
400            access_key_id: "AKIAIOSFODNN7EXAMPLE".to_string(),
401            secret_access_key: "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY".to_string(),
402            session_token: None,
403        }
404    }
405
406    fn hello_context() -> Context {
407        Context::new(vec![ChatMessage::user("Hello")], vec![])
408    }
409
410    #[tokio::test]
411    async fn mantle_disabled_uses_none_without_summary() {
412        let model = LlmModel::all()
413            .iter()
414            .find(|model| {
415                model.provider_enum() == Provider::Bedrock
416                    && model.transport().is_some()
417                    && model.supports_reasoning_off()
418            })
419            .unwrap();
420        let mut server = CaptureServer::start_responses().await;
421        let provider = mantle_provider(&server).await.with_model(&model.model_id());
422        let mut context = hello_context();
423        context.set_reasoning_effort(ReasoningEffort::Disabled);
424        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
425        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
426        let body = server.captured().await.body;
427        assert_eq!(body["reasoning"]["effort"], "none");
428        assert!(body["reasoning"]["summary"].is_null());
429    }
430
431    #[tokio::test]
432    async fn converse_disabled_is_a_non_retryable_transport_error() {
433        let model = LlmModel::all()
434            .iter()
435            .find(|model| {
436                model.provider_enum() == Provider::Bedrock
437                    && model.transport().is_none()
438                    && model.supports_reasoning_off()
439            })
440            .unwrap();
441        let provider =
442            BedrockProvider::new(ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() })
443                .await
444                .with_model(&model.model_id());
445        let mut context = hello_context();
446        context.set_reasoning_effort(crate::ReasoningEffort::Disabled);
447        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
448        assert_eq!(responses.len(), 1);
449        assert!(matches!(&responses[0], Err(LlmError::UnsupportedDisableTransport { .. })));
450        assert!(!responses[0].as_ref().unwrap_err().is_retryable());
451    }
452
453    #[tokio::test]
454    async fn responses_shape_models_are_sent_to_the_responses_endpoint() {
455        let mut server = CaptureServer::start_responses().await;
456        let provider = mantle_provider(&server).await;
457        let mut context = hello_context();
458        context.set_reasoning_effort(crate::ReasoningEffort::Xhigh);
459
460        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
461        let captured = server.captured().await;
462
463        assert!(!responses.is_empty());
464        assert_eq!(captured.path, "/responses");
465        assert_eq!(captured.body["model"], mantle_model());
466        assert_eq!(captured.body["stream"], true);
467        assert_eq!(captured.body["store"], false);
468        assert_eq!(captured.body["reasoning"]["effort"], "xhigh");
469        assert_eq!(captured.body["input"][0]["role"], "user");
470        assert_eq!(captured.body["input"][0]["content"][0]["type"], "input_text");
471        assert_eq!(captured.body["input"][0]["content"][0]["text"], "Hello");
472        assert!(captured.headers.get("authorization").is_none(), "{:?}", captured.headers);
473    }
474
475    #[tokio::test]
476    async fn http_200_failed_server_error_is_retryable_with_diagnostics() {
477        use crate::providers::test_capture_server::ResponseSpec;
478        let spec = ResponseSpec::sse(include_str!("../../../tests/fixtures/openai_responses/04_failed_server.sse"))
479            .with_header("x-amzn-requestid", "amzn-req-123");
480        let mut server = CaptureServer::start_with_spec(spec).await;
481        let provider = mantle_provider(&server).await;
482
483        let responses = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
484        let _ = server.captured().await;
485
486        assert!(!responses.iter().any(|r| matches!(r, Ok(crate::LlmResponse::Done { .. }))));
487        let err = responses.iter().find_map(|r| r.as_ref().err()).expect("expected a failure");
488        assert!(err.is_retryable(), "server_error must be retryable: {err:?}");
489        let provider_error = err.provider().expect("expected provider error");
490        assert_eq!(provider_error.kind, crate::ProviderErrorKind::Server);
491        assert_eq!(provider_error.code.as_deref(), Some("server_error"));
492        assert_eq!(provider_error.http_status, Some(200));
493        assert_eq!(provider_error.request_id.as_deref(), Some("amzn-req-123"));
494    }
495
496    #[tokio::test]
497    async fn failed_event_with_unknown_code_is_terminal_without_done() {
498        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";
499        let mut server = CaptureServer::start_with_response(body).await;
500        let provider = mantle_provider(&server).await;
501
502        let responses = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
503        let _ = server.captured().await;
504
505        assert!(!responses.iter().any(|r| matches!(r, Ok(crate::LlmResponse::Done { .. }))));
506        let err = responses.iter().find_map(|r| r.as_ref().err()).expect("expected a failure");
507        assert!(!err.is_retryable(), "invalid_prompt must be terminal: {err:?}");
508    }
509
510    #[tokio::test]
511    async fn responses_transport_rejects_malformed_and_truncated_streams() {
512        for response in [
513            "data: {not-json}\n\n",
514            "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_1\"}}\n\ndata: [DONE]\n\n",
515        ] {
516            let mut server = CaptureServer::start_with_response(response).await;
517            let provider = mantle_provider(&server).await;
518
519            let responses = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
520            let _ = server.captured().await;
521
522            assert!(responses.iter().any(|response| {
523                response.as_ref().err().and_then(LlmError::provider).map(|provider| provider.kind)
524                    == Some(crate::ProviderErrorKind::StreamInterrupted)
525            }));
526            assert!(!responses.iter().any(|response| matches!(response, Ok(crate::LlmResponse::Done { .. }))));
527        }
528    }
529
530    #[tokio::test]
531    async fn responses_transport_drops_encrypted_reasoning_from_another_model() {
532        let mut server = CaptureServer::start_responses().await;
533        let provider = mantle_provider(&server).await;
534        let context = Context::new(
535            vec![ChatMessage::Assistant {
536                message_id: MessageId::new(),
537                content: "previous answer".to_string(),
538                reasoning: AssistantReasoning {
539                    summary_text: None,
540                    encrypted_content: Some(EncryptedReasoningContent {
541                        id: "reasoning-id".to_string(),
542                        model: "bedrock:openai.gpt-5.5".parse().unwrap(),
543                        content: "opaque-for-another-model".to_string(),
544                    }),
545                },
546                timestamp: IsoString::now(),
547                tool_calls: vec![],
548            }],
549            vec![],
550        );
551
552        let _ = provider.stream_response(&context).collect::<Vec<_>>().await;
553        let captured = server.captured().await;
554
555        assert!(
556            captured.body["input"].as_array().unwrap().iter().all(|item| item["type"] != "reasoning"),
557            "{}",
558            captured.body
559        );
560    }
561
562    #[tokio::test]
563    async fn converse_shape_models_do_not_use_the_responses_endpoint() {
564        let endpoint = FakeBedrockEndpoint::start().await;
565        let provider = BedrockProvider::new(ProviderConnectionConfig {
566            base_url: Some(endpoint.url.clone()),
567            auth_mode: ProviderAuthMode::None,
568            ..Default::default()
569        })
570        .await
571        .with_model(DEFAULT_MODEL);
572
573        let _ = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
574        let request = endpoint.request.await.expect("fake Bedrock endpoint received no request");
575
576        assert!(request.path.starts_with("/model/"), "{}", request.path);
577    }
578
579    #[tokio::test]
580    async fn bearer_token_authenticates_responses_requests() {
581        let mut server = CaptureServer::start_responses().await;
582        let provider = BedrockProvider::from_config(
583            None,
584            Some("us-west-2"),
585            ProviderConnectionConfig { base_url: Some(server.base_url.clone()), ..Default::default() },
586        )
587        .with_bearer_token("test-token")
588        .with_model(&mantle_model());
589
590        let _ = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
591        let captured = server.captured().await;
592
593        assert_eq!(captured.headers.get("authorization").unwrap(), "Bearer test-token");
594    }
595
596    #[tokio::test]
597    async fn credential_chain_sigv4_signs_responses_requests() {
598        let mut server = CaptureServer::start_responses().await;
599        let provider = BedrockProvider::from_config(
600            Some(static_credentials()),
601            Some("us-west-2"),
602            ProviderConnectionConfig { base_url: Some(server.base_url.clone()), ..Default::default() },
603        )
604        .with_model(&mantle_model());
605
606        let _ = provider.stream_response(&hello_context()).collect::<Vec<_>>().await;
607        let captured = server.captured().await;
608
609        let authorization = captured.headers.get("authorization").expect("request was not signed").to_str().unwrap();
610        assert!(
611            authorization.starts_with("AWS4-HMAC-SHA256 Credential=AKIAIOSFODNN7EXAMPLE/"),
612            "unexpected authorization header: {authorization}"
613        );
614        assert!(authorization.contains("/us-west-2/bedrock/aws4_request"), "{authorization}");
615        assert!(captured.headers.contains_key("x-amz-date"), "{:?}", captured.headers);
616    }
617
618    #[test]
619    fn only_responses_shape_models_route_to_the_mantle_transport() {
620        let mantle = test_provider().with_model(&mantle_model());
621        let converse = test_provider().with_model(DEFAULT_MODEL);
622        let profile = test_provider().with_model("us.anthropic.claude-future-model-v99:0");
623
624        assert!(matches!(mantle.mantle_transport(), Some(ModelTransport::OpenAiResponses { .. })));
625        assert_eq!(converse.mantle_transport(), None);
626        assert_eq!(profile.mantle_transport(), None);
627    }
628
629    #[test]
630    fn explicit_connection_preserves_inference_profile() {
631        let provider = BedrockProvider::from_config(
632            None,
633            Some("us-west-2"),
634            ProviderConnectionConfig { inference_profile_arn: Some("arn:test".to_string()), ..Default::default() },
635        );
636
637        assert_eq!(provider.inference_profile_arn.as_deref(), Some("arn:test"));
638    }
639
640    #[test]
641    fn test_from_config_with_credentials() {
642        let provider =
643            BedrockProvider::from_config(Some(static_credentials()), None, ProviderConnectionConfig::default());
644        assert_eq!(provider.model, DEFAULT_MODEL);
645    }
646
647    #[test]
648    fn test_from_config_with_credentials_and_region() {
649        let credentials =
650            AwsCredentials { session_token: Some("FwoGZXIvYXdzEBYaD...".to_string()), ..static_credentials() };
651
652        let provider =
653            BedrockProvider::from_config(Some(credentials), Some("us-west-2"), ProviderConnectionConfig::default())
654                .with_model("anthropic.claude-opus-4-20250514-v1:0");
655
656        assert_eq!(provider.model, "anthropic.claude-opus-4-20250514-v1:0");
657    }
658
659    #[test]
660    fn test_from_config_with_region_only() {
661        let provider = BedrockProvider::from_config(None, Some("eu-west-1"), ProviderConnectionConfig::default());
662        assert_eq!(provider.model, DEFAULT_MODEL);
663    }
664
665    #[test]
666    fn catalog_foundation_id_resolves_context_window() {
667        let provider = test_provider().with_model("anthropic.claude-sonnet-4-5-20250929-v1:0");
668        assert!(provider.context_window().is_some());
669        assert_eq!(provider.model().unwrap().to_string(), "bedrock:anthropic.claude-sonnet-4-5-20250929-v1:0");
670    }
671
672    #[test]
673    fn cross_region_profile_id_in_catalog_resolves() {
674        let provider = test_provider().with_model("us.anthropic.claude-opus-4-6-v1");
675        assert!(provider.context_window().is_some());
676    }
677
678    #[test]
679    fn unknown_cross_region_profile_id_falls_through_to_profile() {
680        let id = "us.anthropic.claude-future-model-v99:0";
681        let provider = test_provider().with_model(id);
682        assert_eq!(provider.context_window(), None);
683        assert_eq!(provider.model().unwrap().to_string(), format!("bedrock:{id}"));
684        assert_eq!(provider.display_name(), format!("Bedrock ({id})"));
685    }
686
687    #[tokio::test]
688    async fn separate_inference_profile_arn_is_used_as_request_model_id() {
689        let endpoint = FakeBedrockEndpoint::start().await;
690        let provider = BedrockProvider::new(ProviderConnectionConfig {
691            base_url: Some(endpoint.url.clone()),
692            auth_mode: ProviderAuthMode::None,
693            request_model: None,
694            inference_profile_arn: Some(application_inference_profile_arn().to_string()),
695        })
696        .await
697        .with_model(DEFAULT_MODEL);
698
699        let result = provider.send_converse_stream(&hello_context()).await;
700        let request = endpoint.request.await.expect("fake Bedrock endpoint received no request");
701
702        assert!(result.is_err());
703        assert!(
704            request.path.contains("arn%3Aaws%3Abedrock%3Aus-west-2%3A000000000000%3Aapplication-inference-profile"),
705            "{}",
706            request.path
707        );
708        assert_eq!(provider.context_window(), Some(200_000));
709        assert_eq!(provider.model().unwrap().to_string(), "bedrock:anthropic.claude-sonnet-4-5-20250929-v1:0");
710    }
711
712    #[test]
713    fn with_inference_profile_arn_keeps_canonical_model_identity() {
714        let arn = inference_profile_arn("us.anthropic.claude-sonnet-4-5-20250929-v1:0");
715        let provider =
716            test_provider().with_model("anthropic.claude-sonnet-4-5-20250929-v1:0").with_inference_profile_arn(&arn);
717
718        assert_eq!(provider.request_model_id(), arn);
719        assert_eq!(provider.context_window(), Some(200_000));
720        assert_eq!(provider.model().unwrap().to_string(), "bedrock:anthropic.claude-sonnet-4-5-20250929-v1:0");
721    }
722
723    #[test]
724    fn prompt_caching_support_comes_from_canonical_model() {
725        let cached = test_provider().with_model("anthropic.claude-sonnet-4-5-20250929-v1:0");
726        assert!(cached.model().unwrap().supports_prompt_caching());
727
728        let unknown_profile = test_provider().with_model("us.anthropic.claude-future-model-v99:0");
729        assert!(!unknown_profile.model().unwrap().supports_prompt_caching());
730    }
731
732    struct FakeBedrockEndpoint {
733        url: String,
734        request: oneshot::Receiver<CapturedRequest>,
735    }
736
737    struct CapturedRequest {
738        method: Method,
739        path: String,
740        headers: HeaderMap,
741    }
742
743    #[derive(Clone)]
744    struct FakeBedrockState {
745        request_tx: Arc<Mutex<Option<oneshot::Sender<CapturedRequest>>>>,
746        shutdown_tx: Arc<Mutex<Option<oneshot::Sender<()>>>>,
747    }
748
749    impl FakeBedrockEndpoint {
750        async fn start() -> Self {
751            let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind fake Bedrock endpoint");
752            let url = format!("http://{}", listener.local_addr().expect("fake Bedrock endpoint address"));
753            let (request_tx, request) = oneshot::channel();
754            let (shutdown_tx, shutdown) = oneshot::channel();
755            let state = FakeBedrockState {
756                request_tx: Arc::new(Mutex::new(Some(request_tx))),
757                shutdown_tx: Arc::new(Mutex::new(Some(shutdown_tx))),
758            };
759
760            let app = Router::new().fallback(any(capture_bedrock_request)).with_state(state);
761            tokio::spawn(async move {
762                axum::serve(listener, app)
763                    .with_graceful_shutdown(async {
764                        let _ = shutdown.await;
765                    })
766                    .await
767                    .expect("serve fake Bedrock endpoint");
768            });
769
770            Self { url, request }
771        }
772    }
773
774    async fn capture_bedrock_request(
775        State(state): State<FakeBedrockState>,
776        request: Request<Body>,
777    ) -> impl IntoResponse {
778        let (parts, _) = request.into_parts();
779        if let Some(tx) = state.request_tx.lock().await.take() {
780            let _ = tx.send(CapturedRequest {
781                method: parts.method,
782                path: parts.uri.path().to_string(),
783                headers: parts.headers,
784            });
785        }
786        if let Some(tx) = state.shutdown_tx.lock().await.take() {
787            let _ = tx.send(());
788        }
789        (StatusCode::FORBIDDEN, "{}")
790    }
791}