1use std::borrow::Cow;
2use std::str::FromStr;
3
4use crate::ReasoningEffort;
5use crate::catalog::transport::ModelTransport;
6use crate::catalog::{BedrockFoundationModel, ModelPricing};
7
8#[derive(Debug, Clone, PartialEq, Eq, Hash)]
9pub enum BedrockModel {
10 Foundation(BedrockFoundationModel),
11 Profile(String),
12}
13
14impl BedrockModel {
15 pub fn model_id(&self) -> Cow<'static, str> {
16 match self {
17 Self::Foundation(m) => Cow::Borrowed(m.model_id()),
18 Self::Profile(s) => Cow::Owned(s.clone()),
19 }
20 }
21
22 pub fn display_name(&self) -> Cow<'static, str> {
23 match self {
24 Self::Foundation(m) => Cow::Borrowed(m.display_name()),
25 Self::Profile(s) => Cow::Owned(format!("Bedrock {s}")),
26 }
27 }
28
29 pub fn context_window(&self) -> Option<u32> {
30 match self {
31 Self::Foundation(m) => Some(m.context_window()),
32 Self::Profile(_) => None,
33 }
34 }
35
36 pub fn reasoning_levels(&self) -> &'static [ReasoningEffort] {
37 match self {
38 Self::Foundation(m) => m.reasoning_levels(),
39 Self::Profile(_) => &[],
40 }
41 }
42
43 pub fn supports_reasoning(&self) -> bool {
44 !self.reasoning_levels().is_empty()
45 }
46
47 pub fn supports_prompt_caching(&self) -> bool {
48 match self {
49 Self::Foundation(m) => m.supports_prompt_caching(),
50 Self::Profile(_) => false,
51 }
52 }
53
54 pub fn supports_image(&self) -> bool {
55 match self {
56 Self::Foundation(m) => m.supports_image(),
57 Self::Profile(_) => false,
58 }
59 }
60
61 pub fn supports_audio(&self) -> bool {
62 match self {
63 Self::Foundation(m) => m.supports_audio(),
64 Self::Profile(_) => false,
65 }
66 }
67
68 pub fn pricing(&self) -> Option<ModelPricing> {
69 match self {
70 Self::Foundation(m) => m.pricing(),
71 Self::Profile(_) => None,
72 }
73 }
74
75 pub fn transport(&self) -> Option<ModelTransport> {
76 match self {
77 Self::Foundation(m) => m.transport(),
78 Self::Profile(_) => None,
79 }
80 }
81}
82
83impl FromStr for BedrockModel {
84 type Err = String;
85
86 fn from_str(s: &str) -> Result<Self, Self::Err> {
87 match s.parse::<BedrockFoundationModel>() {
88 Ok(m) => Ok(Self::Foundation(m)),
89 Err(_) if is_bedrock_inference_profile_arn(s) => Err(
90 "Bedrock inference profile ARNs must be configured as providers.bedrock.inferenceProfileArn; keep model as bedrock:<model-id>".to_string(),
91 ),
92 Err(_) => Ok(Self::Profile(s.to_string())),
93 }
94 }
95}
96
97fn is_bedrock_inference_profile_arn(s: &str) -> bool {
98 let Some(rest) = s.strip_prefix("arn:") else {
99 return false;
100 };
101 let parts: Vec<&str> = rest.split(':').collect();
102 matches!(
103 parts.as_slice(),
104 [partition, "bedrock", _, _, resource, ..]
105 if partition.starts_with("aws")
106 && (resource.starts_with("inference-profile/")
107 || resource.starts_with("application-inference-profile/"))
108 )
109}
110
111#[cfg(test)]
112mod tests {
113 use super::*;
114
115 #[test]
116 fn foundation_model_parses() {
117 let model: BedrockModel = "anthropic.claude-sonnet-4-5-20250929-v1:0".parse().unwrap();
118 assert!(matches!(model, BedrockModel::Foundation(_)));
119 }
120
121 #[test]
122 fn unknown_profile_id_falls_through_to_profile_variant() {
123 let model: BedrockModel = "us.anthropic.claude-future-model-v99:0".parse().unwrap();
124 assert!(matches!(model, BedrockModel::Profile(_)));
125 assert_eq!(model.context_window(), None);
126 }
127
128 #[test]
129 fn inference_profile_arn_is_rejected() {
130 let error =
131 "arn:aws:bedrock:us-west-2:000000000000:inference-profile/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
132 .parse::<BedrockModel>()
133 .unwrap_err();
134 assert!(error.contains("providers.bedrock.inferenceProfileArn"));
135 }
136
137 #[test]
138 fn application_inference_profile_arn_is_rejected() {
139 let error = "arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000"
140 .parse::<BedrockModel>()
141 .unwrap_err();
142 assert!(error.contains("providers.bedrock.inferenceProfileArn"));
143 }
144
145 #[test]
146 fn gov_cloud_arn_is_rejected() {
147 let error = "arn:aws-us-gov:bedrock:us-gov-west-1:000000000000:application-inference-profile/000000000000"
148 .parse::<BedrockModel>()
149 .unwrap_err();
150 assert!(error.contains("providers.bedrock.inferenceProfileArn"));
151 }
152
153 #[test]
154 fn non_bedrock_arn_falls_through_to_profile() {
155 let model: BedrockModel = "arn:aws:s3:us-west-2:000000000000:bucket/foo".parse().unwrap();
156 assert!(matches!(model, BedrockModel::Profile(_)));
157 }
158}