Skip to main content

aether_cli/
provider_connection_args.rs

1use std::collections::BTreeMap;
2use std::str::FromStr;
3
4use llm::{ProviderAuthMode, ProviderConnectionOverride, ProviderConnectionOverrides};
5
6#[derive(Clone, Debug, Default, clap::Args)]
7pub struct ProviderConnectionArgs {
8    #[arg(
9        long = "provider",
10        value_name = "PROVIDER.url=URL|PROVIDER.auth=default|none|PROVIDER.request-model=MODEL|bedrock.inference-profile-arn=ARN"
11    )]
12    pub providers: Vec<ProviderArg>,
13}
14
15impl ProviderConnectionArgs {
16    pub fn into_overrides(self) -> ProviderConnectionOverrides {
17        let mut providers = BTreeMap::new();
18        for arg in self.providers {
19            providers.entry(arg.provider).or_insert_with(ProviderConnectionOverride::default).merge(arg.connection);
20        }
21        ProviderConnectionOverrides::new(providers)
22    }
23}
24
25#[derive(Clone, Debug, PartialEq, Eq)]
26pub struct ProviderArg {
27    provider: String,
28    connection: ProviderConnectionOverride,
29}
30
31impl FromStr for ProviderArg {
32    type Err = String;
33
34    fn from_str(value: &str) -> Result<Self, Self::Err> {
35        let (key, setting) = split_key_value(value)?;
36        let (provider, field) = key
37            .split_once('.')
38            .ok_or_else(|| "provider override must be PROVIDER.url=URL, PROVIDER.auth=default|none, PROVIDER.request-model=MODEL, or bedrock.inference-profile-arn=ARN".to_string())?;
39
40        validate_provider(provider)?;
41        if setting.trim().is_empty() {
42            return Err("provider value cannot be empty".to_string());
43        }
44
45        let connection = match field {
46            "url" => {
47                validate_url(setting)?;
48                ProviderConnectionOverride::url(setting)
49            }
50            "auth" => ProviderConnectionOverride::auth(parse_auth_mode(setting)?),
51            "request-model" => ProviderConnectionOverride::request_model(setting),
52            "inference-profile-arn" => {
53                if provider != "bedrock" {
54                    return Err("inference-profile-arn is only supported for the bedrock provider".to_string());
55                }
56                ProviderConnectionOverride::inference_profile_arn(setting)
57            }
58            _ => {
59                return Err(
60                    "provider override field must be url, auth, request-model, or inference-profile-arn".to_string()
61                );
62            }
63        };
64
65        Ok(Self { provider: provider.to_string(), connection })
66    }
67}
68
69fn split_key_value(value: &str) -> Result<(&str, &str), String> {
70    value.split_once('=').ok_or_else(|| "provider override must be PROVIDER.FIELD=VALUE".to_string())
71}
72
73fn validate_provider(provider: &str) -> Result<(), String> {
74    if provider.trim().is_empty() {
75        return Err("provider name cannot be empty".to_string());
76    }
77    Ok(())
78}
79
80fn validate_url(url: &str) -> Result<(), String> {
81    let parsed = url::Url::parse(url).map_err(|error| format!("invalid provider URL: {error}"))?;
82    match parsed.scheme() {
83        "http" | "https" => Ok(()),
84        scheme => Err(format!("provider URL must use http or https, got {scheme}")),
85    }
86}
87
88fn parse_auth_mode(value: &str) -> Result<ProviderAuthMode, String> {
89    match value {
90        "default" => Ok(ProviderAuthMode::Default),
91        "none" => Ok(ProviderAuthMode::None),
92        _ => Err("provider auth mode must be default or none".to_string()),
93    }
94}
95
96#[cfg(test)]
97mod tests {
98    use super::*;
99
100    #[test]
101    fn parses_provider_url() {
102        let arg: ProviderArg = "bedrock.url=http://127.0.0.1:8787".parse().unwrap();
103        assert_eq!(arg.provider, "bedrock");
104        assert_eq!(arg.connection.base_url.as_deref(), Some("http://127.0.0.1:8787"));
105    }
106
107    #[test]
108    fn parses_provider_auth_modes() {
109        assert_eq!(
110            "bedrock.auth=none".parse::<ProviderArg>().unwrap().connection.auth_mode,
111            Some(ProviderAuthMode::None)
112        );
113        assert_eq!(
114            "bedrock.auth=default".parse::<ProviderArg>().unwrap().connection.auth_mode,
115            Some(ProviderAuthMode::Default)
116        );
117    }
118
119    #[test]
120    fn parses_provider_request_model() {
121        let arg: ProviderArg = "azure-foundry.request-model=production-coding".parse().unwrap();
122        assert_eq!(arg.connection.request_model.as_deref(), Some("production-coding"));
123    }
124
125    #[test]
126    fn combines_repeated_provider_overrides() {
127        let args = ProviderConnectionArgs {
128            providers: vec![
129                "bedrock.url=http://127.0.0.1:8787".parse().unwrap(),
130                "bedrock.auth=none".parse().unwrap(),
131                "azure-foundry.request-model=production-coding".parse().unwrap(),
132                "bedrock.inference-profile-arn=arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000"
133                    .parse()
134                    .unwrap(),
135            ],
136        };
137
138        let overrides = args.into_overrides();
139        let config = overrides.config_for("bedrock");
140
141        assert_eq!(config.base_url.as_deref(), Some("http://127.0.0.1:8787"));
142        assert_eq!(config.auth_mode, ProviderAuthMode::None);
143        assert_eq!(overrides.config_for("azure-foundry").request_model.as_deref(), Some("production-coding"));
144        assert_eq!(
145            config.inference_profile_arn.as_deref(),
146            Some("arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000")
147        );
148    }
149
150    #[test]
151    fn parses_bedrock_inference_profile_arn() {
152        let arg: ProviderArg =
153            "bedrock.inference-profile-arn=arn:aws:bedrock:us-west-2:000000000000:inference-profile/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
154                .parse()
155                .unwrap();
156
157        assert_eq!(arg.provider, "bedrock");
158        assert_eq!(
159            arg.connection.inference_profile_arn.as_deref(),
160            Some(
161                "arn:aws:bedrock:us-west-2:000000000000:inference-profile/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
162            )
163        );
164    }
165
166    #[test]
167    fn rejects_invalid_values() {
168        assert!("bedrock".parse::<ProviderArg>().is_err());
169        assert!("bedrock.url".parse::<ProviderArg>().is_err());
170        assert!(".url=http://127.0.0.1:8787".parse::<ProviderArg>().is_err());
171        assert!("bedrock.url=".parse::<ProviderArg>().is_err());
172        assert!("bedrock.url=file:///tmp/proxy".parse::<ProviderArg>().is_err());
173        assert!("bedrock.auth=disabled".parse::<ProviderArg>().is_err());
174        assert!("bedrock.region=us-west-2".parse::<ProviderArg>().is_err());
175        assert!("openai.inference-profile-arn=arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000".parse::<ProviderArg>().is_err());
176    }
177}