Skip to main content

atman_runtime/
config_provider.rs

1use std::sync::Arc;
2
3use crate::model_registry::ProviderEntry;
4use crate::provider::{Provider, ProviderRegistry};
5use crate::providers::anthropic::AnthropicProvider;
6use crate::providers::openai::OpenAiProvider;
7
8/// Whether a config-backed provider can be used by the current process.
9#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10#[non_exhaustive]
11pub enum ConfigProviderAvailability {
12    Available,
13    Disabled,
14    MissingCredential,
15    UnsupportedKind,
16}
17
18pub(crate) fn reconcile_config_provider_deferred(
19    registry: &ProviderRegistry,
20    name: &str,
21    entry: &ProviderEntry,
22) -> (ConfigProviderAvailability, Option<Arc<dyn Provider>>) {
23    reconcile_config_provider_with(registry, name, entry, |name| std::env::var(name).ok())
24}
25
26/// Build a config-backed provider with the same credential and endpoint resolution as live registration.
27pub fn build_config_provider(
28    name: &str,
29    entry: &ProviderEntry,
30) -> Result<Arc<dyn Provider>, ConfigProviderAvailability> {
31    build_config_provider_with(name, entry, |name| std::env::var(name).ok())
32}
33
34/// Inspect config-backed provider availability without exposing credentials.
35pub fn config_provider_availability(entry: &ProviderEntry) -> ConfigProviderAvailability {
36    match resolve_provider_credential(entry, &|name| std::env::var(name).ok()) {
37        Ok(_) => ConfigProviderAvailability::Available,
38        Err(availability) => availability,
39    }
40}
41
42fn reconcile_config_provider_with(
43    registry: &ProviderRegistry,
44    name: &str,
45    entry: &ProviderEntry,
46    read_env: impl Fn(&str) -> Option<String>,
47) -> (ConfigProviderAvailability, Option<Arc<dyn Provider>>) {
48    let registry_key = format!("config:{name}");
49    let provider = match build_config_provider_with(&registry_key, entry, read_env) {
50        Ok(provider) => provider,
51        Err(availability) => {
52            return (availability, registry.take_named(&registry_key));
53        }
54    };
55    (
56        ConfigProviderAvailability::Available,
57        registry.register_named(registry_key, provider),
58    )
59}
60
61fn build_config_provider_with(
62    registry_key: &str,
63    entry: &ProviderEntry,
64    read_env: impl Fn(&str) -> Option<String>,
65) -> Result<Arc<dyn Provider>, ConfigProviderAvailability> {
66    let api_key = resolve_provider_credential(entry, &read_env)?;
67    let base_url = resolve_base_url(entry, &read_env);
68    match entry.kind.as_str() {
69        "anthropic" => {
70            let mut provider = AnthropicProvider::new(registry_key, api_key);
71            if let Some(base_url) = base_url {
72                provider = provider.with_base_url(&base_url);
73            }
74            if let Some(max_tokens) = entry.max_tokens {
75                provider = provider.with_max_tokens(max_tokens);
76            }
77            Ok(Arc::new(provider))
78        }
79        "openai" | "openai-compat" => {
80            let prompt_cache_key = entry.prompt_cache_key.unwrap_or_else(|| {
81                entry.kind == "openai"
82                    && base_url.as_deref().is_none_or(is_official_openai_base_url)
83            });
84            let mut provider = OpenAiProvider::new(registry_key, api_key)
85                .with_reasoning_format(entry.reasoning_format.unwrap_or_else(|| {
86                    crate::providers::openai::OpenAiReasoningFormat::for_provider_kind(&entry.kind)
87                }))
88                .with_prompt_cache_key(prompt_cache_key);
89            if let Some(base_url) = base_url {
90                provider = provider.with_base_url(&base_url);
91            }
92            if let Some(max_tokens) = entry.max_tokens {
93                provider = provider.with_max_tokens(max_tokens);
94            }
95            Ok(Arc::new(provider))
96        }
97        _ => unreachable!("supported provider kind was checked above"),
98    }
99}
100
101fn is_official_openai_base_url(base_url: &str) -> bool {
102    base_url.trim_end_matches('/') == "https://api.openai.com/v1"
103}
104
105fn resolve_provider_credential(
106    entry: &ProviderEntry,
107    read_env: &impl Fn(&str) -> Option<String>,
108) -> Result<String, ConfigProviderAvailability> {
109    if entry.enabled == Some(false) {
110        return Err(ConfigProviderAvailability::Disabled);
111    }
112    if !matches!(
113        entry.kind.as_str(),
114        "anthropic" | "openai" | "openai-compat"
115    ) {
116        return Err(ConfigProviderAvailability::UnsupportedKind);
117    }
118    resolve_api_key(entry, read_env).ok_or(ConfigProviderAvailability::MissingCredential)
119}
120
121fn resolve_api_key(
122    entry: &ProviderEntry,
123    read_env: &impl Fn(&str) -> Option<String>,
124) -> Option<String> {
125    entry
126        .api_key_env
127        .as_deref()
128        .and_then(read_env)
129        .filter(|value| !value.trim().is_empty())
130        .or_else(|| {
131            entry
132                .api_key
133                .clone()
134                .filter(|value| !value.trim().is_empty())
135        })
136        .or_else(|| {
137            let variable = match entry.kind.as_str() {
138                "openai" | "openai-compat" => "OPENAI_API_KEY",
139                "anthropic" => "ANTHROPIC_API_KEY",
140                _ => return None,
141            };
142            read_env(variable).filter(|value| !value.trim().is_empty())
143        })
144}
145
146fn resolve_base_url(
147    entry: &ProviderEntry,
148    read_env: &impl Fn(&str) -> Option<String>,
149) -> Option<String> {
150    entry
151        .base_url
152        .clone()
153        .filter(|value| !value.trim().is_empty())
154        .or_else(|| {
155            let variable = match entry.kind.as_str() {
156                "openai" | "openai-compat" => "OPENAI_BASE_URL",
157                "anthropic" => "ANTHROPIC_BASE_URL",
158                _ => return None,
159            };
160            read_env(variable).filter(|value| !value.trim().is_empty())
161        })
162}
163
164#[cfg(test)]
165mod tests {
166    use std::collections::HashMap;
167
168    use super::*;
169
170    fn read_from(values: &[(&str, &str)]) -> impl Fn(&str) -> Option<String> {
171        let values = values
172            .iter()
173            .map(|(key, value)| ((*key).to_string(), (*value).to_string()))
174            .collect::<HashMap<_, _>>();
175        move |key| values.get(key).cloned()
176    }
177
178    #[test]
179    fn api_key_resolution_matches_cold_and_live_registration_order() {
180        let mut entry = ProviderEntry {
181            kind: "openai-compat".into(),
182            api_key: Some("inline".into()),
183            api_key_env: Some("CUSTOM_API_KEY".into()),
184            ..Default::default()
185        };
186        assert_eq!(
187            resolve_api_key(
188                &entry,
189                &read_from(&[("CUSTOM_API_KEY", "custom"), ("OPENAI_API_KEY", "fallback")]),
190            )
191            .as_deref(),
192            Some("custom")
193        );
194        assert_eq!(
195            resolve_api_key(&entry, &read_from(&[("OPENAI_API_KEY", "fallback")])).as_deref(),
196            Some("inline")
197        );
198        entry.api_key = None;
199        assert_eq!(
200            resolve_api_key(&entry, &read_from(&[("OPENAI_API_KEY", "fallback")])).as_deref(),
201            Some("fallback")
202        );
203    }
204
205    #[test]
206    fn prompt_cache_key_capability_is_conservative_for_custom_endpoints() {
207        let official = ProviderEntry {
208            name: "official".into(),
209            kind: "openai".into(),
210            api_key: Some("test-key".into()),
211            ..Default::default()
212        };
213        let official_provider = build_config_provider("official", &official).unwrap();
214        assert!(official_provider.capabilities().prompt_cache_key);
215        let disabled = ProviderEntry {
216            prompt_cache_key: Some(false),
217            ..official.clone()
218        };
219        assert!(
220            !build_config_provider("official", &disabled)
221                .unwrap()
222                .capabilities()
223                .prompt_cache_key
224        );
225        let env_routed = build_config_provider_with(
226            "official",
227            &official,
228            read_from(&[("OPENAI_BASE_URL", "https://gateway.example/v1")]),
229        )
230        .unwrap();
231        assert!(!env_routed.capabilities().prompt_cache_key);
232
233        let compatible = ProviderEntry {
234            name: "compatible".into(),
235            kind: "openai-compat".into(),
236            api_key: Some("test-key".into()),
237            ..Default::default()
238        };
239        let compatible_provider = build_config_provider("compatible", &compatible).unwrap();
240        assert!(!compatible_provider.capabilities().prompt_cache_key);
241
242        let custom_official_shape = ProviderEntry {
243            name: "gateway".into(),
244            kind: "openai".into(),
245            api_key: Some("test-key".into()),
246            base_url: Some("https://gateway.example/v1".into()),
247            ..Default::default()
248        };
249        let custom_provider = build_config_provider("gateway", &custom_official_shape).unwrap();
250        assert!(!custom_provider.capabilities().prompt_cache_key);
251
252        let opted_in = ProviderEntry {
253            prompt_cache_key: Some(true),
254            ..custom_official_shape
255        };
256        let opted_in_provider = build_config_provider("gateway", &opted_in).unwrap();
257        assert!(opted_in_provider.capabilities().prompt_cache_key);
258    }
259
260    #[test]
261    fn configured_base_url_precedes_kind_fallback() {
262        let mut entry = ProviderEntry {
263            kind: "anthropic".into(),
264            base_url: Some("https://configured.invalid".into()),
265            ..Default::default()
266        };
267        let read_env = read_from(&[("ANTHROPIC_BASE_URL", "https://fallback.invalid")]);
268        assert_eq!(
269            resolve_base_url(&entry, &read_env).as_deref(),
270            Some("https://configured.invalid")
271        );
272        entry.base_url = None;
273        assert_eq!(
274            resolve_base_url(&entry, &read_env).as_deref(),
275            Some("https://fallback.invalid")
276        );
277    }
278
279    #[test]
280    fn reconciliation_removes_disabled_or_unusable_provider() {
281        let registry = ProviderRegistry::new();
282        let mut entry = ProviderEntry {
283            kind: "openai-compat".into(),
284            api_key: Some("test-key".into()),
285            enabled: Some(true),
286            ..Default::default()
287        };
288        assert_eq!(
289            reconcile_config_provider_with(&registry, "gateway", &entry, read_from(&[])).0,
290            ConfigProviderAvailability::Available
291        );
292        assert!(registry.contains("config:gateway"));
293
294        entry.enabled = Some(false);
295        assert_eq!(
296            reconcile_config_provider_with(&registry, "gateway", &entry, read_from(&[])).0,
297            ConfigProviderAvailability::Disabled
298        );
299        assert!(!registry.contains("config:gateway"));
300
301        entry.enabled = Some(true);
302        entry.api_key = None;
303        assert_eq!(
304            reconcile_config_provider_with(&registry, "gateway", &entry, read_from(&[])).0,
305            ConfigProviderAvailability::MissingCredential
306        );
307        assert!(!registry.contains("config:gateway"));
308    }
309
310    #[test]
311    fn availability_distinguishes_disabled_missing_and_unsupported_providers() {
312        let mut entry = ProviderEntry {
313            kind: "openai-compat".into(),
314            api_key: Some("test-key".into()),
315            enabled: Some(true),
316            ..Default::default()
317        };
318        assert_eq!(
319            build_config_provider_with("config:test", &entry, read_from(&[]))
320                .map(|_| ConfigProviderAvailability::Available)
321                .unwrap_or_else(|availability| availability),
322            ConfigProviderAvailability::Available
323        );
324
325        entry.enabled = Some(false);
326        assert!(matches!(
327            build_config_provider_with("config:test", &entry, read_from(&[])),
328            Err(ConfigProviderAvailability::Disabled)
329        ));
330
331        entry.enabled = Some(true);
332        entry.api_key = None;
333        assert!(matches!(
334            build_config_provider_with("config:test", &entry, read_from(&[])),
335            Err(ConfigProviderAvailability::MissingCredential)
336        ));
337
338        entry.kind = "unsupported".into();
339        assert!(matches!(
340            build_config_provider_with("config:test", &entry, read_from(&[])),
341            Err(ConfigProviderAvailability::UnsupportedKind)
342        ));
343    }
344}