Skip to main content

aether_cli/
provider_connection_args.rs

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