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