1use std::sync::Arc;
9
10use tokio::sync::Mutex;
11
12use mermaid_domain::Config;
13use mermaid_model::models::{ModelError, Result, lookup_provider};
14use mermaid_model::utils::{
15 resolve_api_key, resolve_provider_key, resolve_provider_key_with_fallback,
16};
17
18use super::model::{anthropic, gemini, meta};
19
20fn require_key(provider: &str, default_env: &str, override_env: Option<&str>) -> Result<String> {
24 resolve_provider_key(provider, default_env, override_env).ok_or_else(|| {
25 let env = override_env.unwrap_or(default_env);
26 ModelError::Authentication(format!(
27 "{provider} requires env var {env} (or `mermaid login {provider}`)"
28 ))
29 })
30}
31
32fn require_key_with_fallback(
33 provider: &str,
34 env_var: &str,
35 fallback_env_var: &str,
36) -> Result<String> {
37 resolve_provider_key_with_fallback(provider, env_var, fallback_env_var, None).ok_or_else(|| {
38 ModelError::Authentication(format!(
39 "{provider} requires env var {env_var} (or legacy {fallback_env_var}, or \
40 `mermaid login {provider}`)"
41 ))
42 })
43}
44
45fn key_env_source(default_env: &str, override_env: Option<&str>) -> Option<String> {
51 let env = override_env.unwrap_or(default_env);
55 resolve_api_key(env, None).map(|_| env.to_string())
56}
57
58fn key_env_source_with_fallback(default_env: &str, fallback_env: &str) -> Option<String> {
60 if resolve_api_key(default_env, None).is_some() {
61 return Some(default_env.to_string());
62 }
63 if resolve_api_key(fallback_env, None).is_some() {
64 return Some(fallback_env.to_string());
65 }
66 None
67}
68
69#[derive(Debug, Clone, PartialEq, Eq)]
79pub(crate) struct ProviderEndpoint {
80 pub(crate) base_url: String,
82 pub(crate) api_key: Option<String>,
85 pub(crate) key_env: Option<String>,
88}
89
90impl ProviderEndpoint {
91 fn require_key(&self, provider: &str) -> Result<String> {
97 self.api_key
98 .clone()
99 .ok_or_else(|| ModelError::Authentication(format!("{provider} requires an API key")))
100 }
101}
102
103fn is_known_provider(config: &Config, provider_lc: &str) -> bool {
106 matches!(provider_lc, "anthropic" | "gemini" | "meta")
107 || lookup_provider(provider_lc).is_some()
108 || config.providers.contains_key(provider_lc)
109}
110
111pub(crate) fn resolve_provider_endpoint(
119 config: &Config,
120 provider: &str,
121) -> Result<ProviderEndpoint> {
122 let provider_lc = provider.to_lowercase();
123 let user_cfg = config.providers.get(&provider_lc);
124 let override_env = user_cfg.and_then(|c| c.api_key_env.as_deref());
125 let override_url = user_cfg.and_then(|c| c.base_url.clone());
126
127 let bespoke = match provider_lc.as_str() {
130 "anthropic" => Some((
131 anthropic::DEFAULT_API_KEY_ENV,
132 None,
133 anthropic::DEFAULT_BASE_URL,
134 )),
135 "gemini" => Some((
136 gemini::DEFAULT_API_KEY_ENV,
137 Some(gemini::LEGACY_API_KEY_ENV),
138 gemini::DEFAULT_BASE_URL,
139 )),
140 "meta" => Some((meta::DEFAULT_API_KEY_ENV, None, meta::DEFAULT_BASE_URL)),
141 _ => None,
142 };
143 if let Some((default_env, legacy_env, default_url)) = bespoke {
144 let base_url = resolve_overridable_base_url(&provider_lc, override_url, default_url)?;
145 let (api_key, key_env) = match legacy_env {
146 Some(legacy) if override_env.is_none() => (
149 require_key_with_fallback(&provider_lc, default_env, legacy)?,
150 key_env_source_with_fallback(default_env, legacy),
151 ),
152 _ => (
153 require_key(&provider_lc, default_env, override_env)?,
154 key_env_source(default_env, override_env),
155 ),
156 };
157 return Ok(ProviderEndpoint {
158 base_url,
159 api_key: Some(api_key),
160 key_env,
161 });
162 }
163
164 if provider_lc == "cloudflare" {
170 return resolve_cloudflare_endpoint(&provider_lc, override_env, override_url);
171 }
172
173 if let Some(profile) = lookup_provider(&provider_lc) {
175 let base_url = resolve_overridable_base_url(&provider_lc, override_url, profile.base_url)?;
176 let api_key = resolve_optional_key(
180 &provider_lc,
181 profile.api_key_env,
182 override_env,
183 &base_url,
184 profile.key_hint,
185 )?;
186 let key_env = api_key
187 .is_some()
188 .then(|| key_env_source(profile.api_key_env, override_env))
189 .flatten();
190 return Ok(ProviderEndpoint {
191 base_url,
192 api_key,
193 key_env,
194 });
195 }
196
197 if user_cfg.is_some() {
200 return resolve_custom_endpoint(&provider_lc, override_env, override_url);
201 }
202
203 Err(ModelError::InvalidRequest(format!(
204 "Unknown provider '{provider}'"
205 )))
206}
207
208fn resolve_cloudflare_endpoint(
213 provider_lc: &str,
214 override_env: Option<&str>,
215 override_url: Option<String>,
216) -> Result<ProviderEndpoint> {
217 let profile = lookup_provider("cloudflare").expect("cloudflare is in the registry");
218 let api_key_env = override_env.unwrap_or(profile.api_key_env);
219 let base_url = match override_url {
220 Some(url) => {
223 validate_provider_base_url(&url)?;
224 warn_overridden_provider_host("cloudflare", &url);
225 url
226 },
227 None => match require_cloudflare_account_id() {
231 Ok(id) => cloudflare_base_url(&id),
232 Err(_)
233 if resolve_provider_key("cloudflare", profile.api_key_env, override_env)
234 .is_none() =>
235 {
236 return Err(ModelError::Authentication(format!(
237 "cloudflare requires env vars {api_key_env} and CLOUDFLARE_ACCOUNT_ID — \
238 create a token at https://dash.cloudflare.com/profile/api-tokens; the \
239 account id is on your Cloudflare dashboard (or set \
240 [providers.cloudflare].base_url)"
241 )));
242 },
243 Err(e) => return Err(e),
244 },
245 };
246 let api_key = resolve_optional_key(
247 provider_lc,
248 profile.api_key_env,
249 override_env,
250 &base_url,
251 profile.key_hint,
252 )?;
253 let key_env = api_key
254 .is_some()
255 .then(|| key_env_source(profile.api_key_env, override_env))
256 .flatten();
257 Ok(ProviderEndpoint {
258 base_url,
259 api_key,
260 key_env,
261 })
262}
263
264fn resolve_custom_endpoint(
267 provider_lc: &str,
268 override_env: Option<&str>,
269 override_url: Option<String>,
270) -> Result<ProviderEndpoint> {
271 let base_url = override_url.ok_or_else(|| {
272 ModelError::InvalidRequest(format!(
273 "custom provider '{provider_lc}' requires base_url in config"
274 ))
275 })?;
276 let resolved = match override_env {
283 Some(env) => resolve_provider_key(provider_lc, env, None),
284 None => mermaid_model::utils::default_store().get(provider_lc),
285 };
286 let (api_key, key_env) = match resolved {
287 Some(key) => {
288 validate_provider_base_url(&base_url)?;
289 let env = override_env.and_then(|env| key_env_source(env, None));
290 (Some(key), env)
291 },
292 None if base_url_is_local(&base_url) => (None, None),
293 None => {
294 let reason = match override_env {
295 Some(env) => format!(
296 "requires env var {env} (or `mermaid login {provider_lc}`, or a \
297 loopback/LAN base_url)"
298 ),
299 None => "requires api_key_env, or a loopback/LAN base_url".to_string(),
300 };
301 return Err(ModelError::Authentication(format!(
302 "custom provider '{provider_lc}' {reason}"
303 )));
304 },
305 };
306 Ok(ProviderEndpoint {
307 base_url,
308 api_key,
309 key_env,
310 })
311}
312
313fn resolve_optional_key(
318 provider: &str,
319 default_env: &str,
320 override_env: Option<&str>,
321 base_url: &str,
322 hint: Option<&str>,
323) -> Result<Option<String>> {
324 if let Some(key) = resolve_provider_key(provider, default_env, override_env) {
325 return Ok(Some(key));
326 }
327 if base_url_is_local(base_url) {
328 return Ok(None);
329 }
330 let env = override_env.unwrap_or(default_env);
331 let mut msg = format!("{provider} requires env var {env} (or `mermaid login {provider}`)");
332 if let Some(h) = hint {
333 msg.push_str(" — ");
334 msg.push_str(h);
335 }
336 Err(ModelError::Authentication(msg))
337}
338
339fn base_url_is_local(base_url: &str) -> bool {
342 reqwest::Url::parse(base_url)
343 .ok()
344 .and_then(|u| {
345 u.host_str()
346 .map(|h| mermaid_model::utils::classify_host(h).is_internal())
347 })
348 .unwrap_or(false)
349}
350
351fn merged_headers(
357 profile: &mermaid_model::models::ProviderProfile,
358 user_cfg: Option<&mermaid_domain::UserProviderConfig>,
359) -> std::collections::HashMap<String, String> {
360 let mut headers: std::collections::HashMap<String, String> = profile
361 .extra_headers
362 .iter()
363 .map(|(k, v)| ((*k).to_string(), (*v).to_string()))
364 .collect();
365 if let Some(cfg) = user_cfg {
366 for (k, v) in &cfg.extra_headers {
367 headers.insert(k.clone(), v.clone());
368 }
369 for (header, env_var) in &cfg.env_headers {
370 if let Ok(val) = std::env::var(env_var) {
371 headers.insert(header.clone(), val);
372 }
373 }
374 }
375 headers
376}
377
378use super::model::{
379 AnthropicProvider, GeminiProvider, MetaProvider, ModelProvider, OllamaProvider,
380 OpenAICompatProvider,
381};
382
383type ProviderCell = Arc<tokio::sync::OnceCell<Arc<dyn ModelProvider>>>;
387
388pub struct ProviderFactory {
392 config: Arc<Config>,
393 cache: Mutex<std::collections::HashMap<String, ProviderCell>>,
396}
397
398impl ProviderFactory {
399 #[must_use]
400 pub fn new(config: Config) -> Self {
401 Self {
402 config: Arc::new(config),
403 cache: Mutex::new(std::collections::HashMap::new()),
404 }
405 }
406
407 pub fn with_seeded_providers(
418 config: Config,
419 seeds: impl IntoIterator<Item = (String, Arc<dyn ModelProvider>)>,
420 ) -> Self {
421 let cache = seeds
422 .into_iter()
423 .map(|(model_id, provider)| {
424 let cell = tokio::sync::OnceCell::new_with(Some(provider));
425 (normalize_cache_key(&model_id), Arc::new(cell))
426 })
427 .collect();
428 Self {
429 config: Arc::new(config),
430 cache: Mutex::new(cache),
431 }
432 }
433
434 pub fn config(&self) -> &Config {
435 &self.config
436 }
437
438 pub async fn resolve(&self, model_id: &str) -> Result<Arc<dyn ModelProvider>> {
449 let key = normalize_cache_key(model_id);
450 let cell = {
454 let mut cache = self.cache.lock().await;
455 Arc::clone(
456 cache
457 .entry(key)
458 .or_insert_with(|| Arc::new(tokio::sync::OnceCell::new())),
459 )
460 };
461 let provider = cell
462 .get_or_try_init(|| async {
463 let p = build_provider(&self.config, model_id).await?;
464 Ok::<Arc<dyn ModelProvider>, ModelError>(Arc::from(p))
465 })
466 .await?;
467 Ok(Arc::clone(provider))
468 }
469}
470
471async fn build_provider(config: &Config, model_id: &str) -> Result<Box<dyn ModelProvider>> {
480 let (provider, model_name) = parse_model_id(model_id);
481 let provider_lc = provider.to_lowercase();
482
483 if provider_lc == "ollama" {
486 let backend = crate::ollama::backend_config(config);
487 let p = OllamaProvider::with_app_config(
488 model_name,
489 Arc::new(backend),
490 Arc::new(config.clone()),
491 )
492 .await?;
493 return Ok(Box::new(p));
494 }
495
496 if !is_known_provider(config, &provider_lc) {
501 return Err(ModelError::InvalidRequest(format!(
502 "Unknown provider '{provider}' (model_id: {model_id})"
503 )));
504 }
505 let endpoint = resolve_provider_endpoint(config, &provider_lc)?;
506 let user_cfg = config.providers.get(&provider_lc);
507
508 if provider_lc == "anthropic" {
510 let p = AnthropicProvider::new(
511 endpoint.require_key(&provider_lc)?,
512 model_name.to_string(),
513 endpoint.base_url,
514 )?;
515 return Ok(Box::new(p));
516 }
517
518 if provider_lc == "gemini" {
520 let p = GeminiProvider::new(
521 endpoint.require_key(&provider_lc)?,
522 model_name.to_string(),
523 endpoint.base_url,
524 )?;
525 return Ok(Box::new(p));
526 }
527
528 if provider_lc == "meta" {
531 let api_key = endpoint.require_key(&provider_lc)?;
532 let mut extra_headers = std::collections::HashMap::new();
533 if let Some(cfg) = user_cfg {
534 extra_headers.extend(cfg.extra_headers.clone());
535 for (header, env_var) in &cfg.env_headers {
536 if let Ok(value) = std::env::var(env_var) {
537 extra_headers.insert(header.clone(), value);
538 }
539 }
540 }
541 let p = MetaProvider::new(
542 api_key,
543 model_name.to_string(),
544 endpoint.base_url,
545 extra_headers,
546 )?;
547 return Ok(Box::new(p));
548 }
549
550 let profile = match lookup_provider(&provider_lc) {
554 Some(profile) => profile,
555 None => user_cfg
558 .and_then(|cfg| user_profile_to_static(&provider_lc, cfg))
559 .ok_or_else(|| {
560 ModelError::InvalidRequest(format!(
561 "Unknown provider '{provider}' (model_id: {model_id})"
562 ))
563 })?,
564 };
565 let extra_headers = merged_headers(profile, user_cfg);
566 let p = OpenAICompatProvider::new(
567 profile,
568 endpoint.base_url,
569 endpoint.api_key,
570 model_name.to_string(),
571 extra_headers,
572 )?;
573 Ok(Box::new(p))
574}
575
576fn normalize_cache_key(model_id: &str) -> String {
580 let (provider, model) = parse_model_id(model_id);
581 format!("{}/{}", provider.to_lowercase(), model)
582}
583
584fn parse_model_id(model_id: &str) -> (String, &str) {
587 match model_id.split_once('/') {
588 Some((p, m)) => (p.to_string(), m),
589 None => ("ollama".to_string(), model_id),
590 }
591}
592
593#[must_use]
605pub fn model_provider_resolves(config: &Config, model_id: &str) -> bool {
606 let (provider, _) = parse_model_id(model_id);
607 provider.eq_ignore_ascii_case("ollama") || resolve_provider_endpoint(config, &provider).is_ok()
608}
609
610static PROFILE_CACHE: std::sync::LazyLock<
614 std::sync::Mutex<
615 std::collections::HashMap<String, &'static mermaid_model::models::ProviderProfile>,
616 >,
617> = std::sync::LazyLock::new(|| std::sync::Mutex::new(std::collections::HashMap::new()));
618
619fn user_profile_to_static(
633 name: &str,
634 user_cfg: &mermaid_domain::UserProviderConfig,
635) -> Option<&'static mermaid_model::models::ProviderProfile> {
636 use mermaid_model::models::{ProviderProfile, ReasoningExtraction, ReasoningStrategy};
637
638 let compat = user_cfg.compat.as_deref().unwrap_or("openai");
639 let base_url = user_cfg.base_url.clone().unwrap_or_default();
640 let api_key_env = user_cfg.api_key_env.clone().unwrap_or_default();
641
642 let cache_key = format!("{name}\u{0}{base_url}\u{0}{api_key_env}\u{0}{compat}");
645
646 let mut cache = PROFILE_CACHE
647 .lock()
648 .unwrap_or_else(|poisoned| poisoned.into_inner());
649 if let Some(profile) = cache.get(cache_key.as_str()) {
650 return Some(*profile);
653 }
654
655 let strategy = match compat {
656 "openai" => ReasoningStrategy::None,
657 "openai-effort" => ReasoningStrategy::Effort,
658 "openrouter" => ReasoningStrategy::OpenRouterShape,
659 _ => ReasoningStrategy::None,
660 };
661
662 let profile = Box::new(ProviderProfile {
663 name: Box::leak(name.to_string().into_boxed_str()),
664 base_url: Box::leak(base_url.into_boxed_str()),
665 api_key_env: Box::leak(api_key_env.into_boxed_str()),
666 key_hint: None,
667 extra_headers: &[],
668 reasoning_strategy: strategy,
669 reasoning_extraction: ReasoningExtraction::None,
670 max_tokens_param: mermaid_model::models::MaxTokensParam::MaxTokens,
671 disable_parallel_tool_calls_for: &[],
672 });
673 let leaked: &'static ProviderProfile = Box::leak(profile);
674 cache.insert(cache_key, leaked);
675 Some(leaked)
676}
677
678fn cloudflare_base_url(account_id: &str) -> String {
683 format!(
684 "https://api.cloudflare.com/client/v4/accounts/{}/ai/v1",
685 account_id.trim()
686 )
687}
688
689pub(crate) fn discovery_base_url(
697 profile: &mermaid_model::models::ProviderProfile,
698 override_url: Option<String>,
699) -> Option<String> {
700 if override_url.is_some() {
701 return override_url;
702 }
703 if profile.name == "cloudflare" {
704 return require_cloudflare_account_id()
705 .ok()
706 .map(|id| cloudflare_base_url(&id));
707 }
708 Some(profile.base_url.to_string())
709}
710
711fn require_cloudflare_account_id() -> Result<String> {
716 resolve_api_key("CLOUDFLARE_ACCOUNT_ID", None)
717 .map(|s| s.trim().to_string())
718 .filter(|s| !s.is_empty())
719 .ok_or_else(|| {
720 ModelError::Authentication(
721 "cloudflare requires env var CLOUDFLARE_ACCOUNT_ID (your Cloudflare account id) — \
722 find it on your Cloudflare dashboard, or set [providers.cloudflare].base_url"
723 .to_string(),
724 )
725 })
726}
727
728fn validate_provider_base_url(url: &str) -> Result<()> {
734 let parsed = reqwest::Url::parse(url).map_err(|e| {
735 ModelError::InvalidRequest(format!("invalid provider base_url '{url}': {e}"))
736 })?;
737 match parsed.scheme() {
738 "https" => Ok(()),
739 "http"
744 if mermaid_model::utils::classify_host(parsed.host_str().unwrap_or_default())
745 .is_loopback() =>
746 {
747 Ok(())
748 },
749 "http" => Err(ModelError::InvalidRequest(format!(
750 "provider base_url '{url}' uses http:// to a non-loopback host — refusing to send the \
751 API key in cleartext. Use https, or http://localhost for a local server."
752 ))),
753 other => Err(ModelError::InvalidRequest(format!(
754 "provider base_url '{url}' has unsupported scheme '{other}' (use http or https)"
755 ))),
756 }
757}
758
759fn resolve_overridable_base_url(
775 provider: &str,
776 override_url: Option<String>,
777 default_url: &str,
778) -> Result<String> {
779 match override_url {
780 Some(url) => {
781 validate_provider_base_url(&url)?;
782 warn_overridden_provider_host(provider, &url);
783 Ok(url)
784 },
785 None => Ok(default_url.to_string()),
786 }
787}
788
789static WARNED_OVERRIDE_HOSTS: std::sync::LazyLock<
793 std::sync::Mutex<std::collections::HashSet<String>>,
794> = std::sync::LazyLock::new(|| std::sync::Mutex::new(std::collections::HashSet::new()));
795
796fn warn_overridden_provider_host(provider: &str, base_url: &str) {
800 let host = provider_host(base_url);
801 if should_warn_once(&format!("{provider}@{host}")) {
802 tracing::warn!(
803 "built-in provider '{}' base_url overridden in config: the {} API key will be sent to \
804 host '{}' instead of the trusted default endpoint",
805 provider,
806 provider,
807 host
808 );
809 }
810}
811
812fn provider_host(base_url: &str) -> String {
815 reqwest::Url::parse(base_url)
816 .ok()
817 .and_then(|u| u.host_str().map(str::to_string))
818 .unwrap_or_else(|| "<unknown>".to_string())
819}
820
821fn should_warn_once(key: &str) -> bool {
824 let mut warned = WARNED_OVERRIDE_HOSTS
825 .lock()
826 .unwrap_or_else(|poisoned| poisoned.into_inner());
827 warned.insert(key.to_string())
828}
829
830#[cfg(test)]
831mod tests {
832 use super::*;
833
834 #[test]
835 fn base_url_is_local_classifies_hosts() {
836 assert!(base_url_is_local("http://127.0.0.1:8000/v1"));
837 assert!(base_url_is_local("http://localhost:1234/v1"));
838 assert!(base_url_is_local("http://192.168.1.5:8000/v1"));
839 assert!(!base_url_is_local("https://api.openai.com/v1"));
840 assert!(!base_url_is_local("not a url"));
841 }
842
843 #[test]
844 fn merged_headers_keeps_static_profile_headers_and_user_overrides() {
845 let profile = mermaid_model::models::lookup_provider("openrouter").unwrap();
846 let base = merged_headers(profile, None);
848 assert_eq!(
849 base.get("X-OpenRouter-Title").map(String::as_str),
850 Some("Mermaid")
851 );
852 assert!(base.contains_key("HTTP-Referer"));
853 let mut cfg = mermaid_domain::UserProviderConfig::default();
855 cfg.extra_headers.insert("X-Custom".into(), "v".into());
856 cfg.extra_headers
857 .insert("X-OpenRouter-Title".into(), "Override".into());
858 let merged = merged_headers(profile, Some(&cfg));
859 assert_eq!(merged.get("X-Custom").map(String::as_str), Some("v"));
860 assert_eq!(
861 merged.get("X-OpenRouter-Title").map(String::as_str),
862 Some("Override")
863 );
864 assert!(merged.contains_key("HTTP-Referer"));
865 }
866
867 #[test]
868 fn merged_headers_resolves_env_headers_and_skips_missing() {
869 let profile = mermaid_model::models::lookup_provider("openai").unwrap();
870 let mut cfg = mermaid_domain::UserProviderConfig::default();
871 cfg.env_headers
872 .insert("X-Gateway-Token".into(), "MERMAID_TEST_GW_TOKEN".into());
873 temp_env::with_var("MERMAID_TEST_GW_TOKEN", Some("secret123"), || {
874 let merged = merged_headers(profile, Some(&cfg));
875 assert_eq!(
876 merged.get("X-Gateway-Token").map(String::as_str),
877 Some("secret123")
878 );
879 });
880 temp_env::with_var("MERMAID_TEST_GW_TOKEN", None::<&str>, || {
881 assert!(!merged_headers(profile, Some(&cfg)).contains_key("X-Gateway-Token"));
882 });
883 }
884
885 #[test]
886 fn base_url_requires_https_for_remote_hosts() {
887 assert!(validate_provider_base_url("http://api.example.com/v1").is_err());
889 assert!(validate_provider_base_url("ftp://example.com").is_err());
890 assert!(validate_provider_base_url("https://api.example.com/v1").is_ok());
892 assert!(validate_provider_base_url("http://localhost:11434/v1").is_ok());
893 assert!(validate_provider_base_url("http://127.0.0.1:8000").is_ok());
894 assert!(validate_provider_base_url("http://[::1]:8000").is_ok());
895 assert!(validate_provider_base_url("http://192.168.1.5:8080").is_err());
898 assert!(validate_provider_base_url("http://169.254.169.254").is_err());
899 }
900
901 #[test]
902 fn cloudflare_base_url_synthesizes_account_scoped_endpoint() {
903 assert_eq!(
904 cloudflare_base_url("acct123"),
905 "https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1"
906 );
907 assert_eq!(
909 cloudflare_base_url(" acct123\n"),
910 "https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1"
911 );
912 }
913
914 #[test]
915 fn cloudflare_account_id_required_and_non_blank() {
916 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", None::<&str>, || {
918 let err = require_cloudflare_account_id().expect_err("must error when unset");
919 assert!(format!("{err}").contains("CLOUDFLARE_ACCOUNT_ID"));
920 });
921 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some(" "), || {
923 assert!(require_cloudflare_account_id().is_err());
924 });
925 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some(" acct123 "), || {
927 assert_eq!(require_cloudflare_account_id().unwrap(), "acct123");
928 });
929 }
930
931 #[test]
932 fn discovery_base_url_resolves_per_provider() {
933 let cf = lookup_provider("cloudflare").expect("cloudflare is in the registry");
934 let openai = lookup_provider("openai").expect("openai is in the registry");
935 assert_eq!(
937 discovery_base_url(cf, Some("https://gw.example/v1".to_string())),
938 Some("https://gw.example/v1".to_string())
939 );
940 assert_eq!(
942 discovery_base_url(openai, None),
943 Some(openai.base_url.to_string())
944 );
945 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some("acct123"), || {
947 assert_eq!(
948 discovery_base_url(cf, None),
949 Some("https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1".to_string())
950 );
951 });
952 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", None::<&str>, || {
954 assert_eq!(discovery_base_url(cf, None), None);
955 });
956 }
957
958 #[tokio::test]
959 async fn cloudflare_missing_both_env_vars_is_one_combined_error() {
960 temp_env::async_with_vars(
961 [
962 ("CLOUDFLARE_ACCOUNT_ID", None::<&str>),
963 ("CLOUDFLARE_API_TOKEN", None),
964 ],
965 async {
966 let f = ProviderFactory::new(Config::default());
967 let err = match f.resolve("cloudflare/@cf/zai-org/glm-5.2").await {
968 Ok(_) => panic!("must fail with neither env var set"),
969 Err(e) => e,
970 };
971 let msg = format!("{err}");
972 assert!(
973 msg.contains("CLOUDFLARE_API_TOKEN") && msg.contains("CLOUDFLARE_ACCOUNT_ID"),
974 "one error must name both missing vars, got: {msg}"
975 );
976 },
977 )
978 .await;
979 }
980
981 use std::sync::atomic::{AtomicUsize, Ordering};
982
983 fn unique_env(prefix: &str) -> String {
984 static N: AtomicUsize = AtomicUsize::new(0);
985 format!(
986 "{}_{}_{}",
987 prefix,
988 std::process::id(),
989 N.fetch_add(1, Ordering::SeqCst)
990 )
991 }
992
993 #[test]
994 fn parse_bare_name_defaults_to_ollama() {
995 let (p, m) = parse_model_id("qwen3-coder:30b");
996 assert_eq!(p, "ollama");
997 assert_eq!(m, "qwen3-coder:30b");
998 }
999
1000 #[test]
1001 fn parse_prefixed() {
1002 let (p, m) = parse_model_id("anthropic/claude-opus-4-7");
1003 assert_eq!(p, "anthropic");
1004 assert_eq!(m, "claude-opus-4-7");
1005 }
1006
1007 #[tokio::test]
1008 async fn meta_requires_its_documented_api_key_env() {
1009 temp_env::async_with_vars(
1010 [(
1011 crate::providers::model::meta::DEFAULT_API_KEY_ENV,
1012 None::<&str>,
1013 )],
1014 async {
1015 let factory = ProviderFactory::new(Config::default());
1016 let error = match factory.resolve("meta/muse-spark-1.1").await {
1017 Ok(_) => panic!("Meta must require an API key"),
1018 Err(error) => error,
1019 };
1020 assert!(
1021 error
1022 .to_string()
1023 .contains(crate::providers::model::meta::DEFAULT_API_KEY_ENV)
1024 );
1025 },
1026 )
1027 .await;
1028 }
1029
1030 #[tokio::test]
1031 async fn meta_routes_to_responses_provider_with_muse_capabilities() {
1032 temp_env::async_with_vars(
1033 [(
1034 crate::providers::model::meta::DEFAULT_API_KEY_ENV,
1035 Some("test-key"),
1036 )],
1037 async {
1038 let factory = ProviderFactory::new(Config::default());
1039 let provider = factory.resolve("meta/muse-spark-1.1").await.unwrap();
1040 let capabilities = provider.capabilities();
1041 assert!(capabilities.supports_tools);
1042 assert!(capabilities.supports_vision);
1043 assert!(capabilities.emits_provider_continuation);
1044 assert_eq!(
1045 capabilities.max_context_tokens,
1046 Some(mermaid_model::constants::META_MUSE_SPARK_CONTEXT_WINDOW)
1047 );
1048 assert_eq!(
1049 capabilities.max_output_tokens,
1050 Some(mermaid_model::constants::META_MUSE_SPARK_MAX_OUTPUT_TOKENS)
1051 );
1052 },
1053 )
1054 .await;
1055 }
1056
1057 #[test]
1058 fn gemini_key_resolution_accepts_legacy_fallback() {
1059 let primary = unique_env("MERMAID_FACTORY_GEMINI_PRIMARY");
1060 let legacy = unique_env("MERMAID_FACTORY_GEMINI_LEGACY");
1061 temp_env::with_vars(
1062 [(primary.as_str(), None), (legacy.as_str(), Some("legacy"))],
1063 || {
1064 let resolved = require_key_with_fallback("gemini", &primary, &legacy)
1065 .expect("legacy fallback should resolve");
1066 assert_eq!(resolved, "legacy");
1067 },
1068 );
1069 }
1070
1071 #[test]
1072 fn gemini_key_resolution_prefers_google_primary() {
1073 let primary = unique_env("MERMAID_FACTORY_GEMINI_PRIMARY2");
1074 let legacy = unique_env("MERMAID_FACTORY_GEMINI_LEGACY2");
1075 temp_env::with_vars(
1076 [
1077 (primary.as_str(), Some("google")),
1078 (legacy.as_str(), Some("legacy")),
1079 ],
1080 || {
1081 let resolved = require_key_with_fallback("gemini", &primary, &legacy)
1082 .expect("primary should resolve");
1083 assert_eq!(resolved, "google");
1084 },
1085 );
1086 }
1087
1088 #[tokio::test]
1089 async fn factory_reports_unknown_provider_clearly() {
1090 let cfg = Config::default();
1091 let f = ProviderFactory::new(cfg);
1092 match f.resolve("totally-made-up/model").await {
1093 Ok(_) => panic!("expected error"),
1094 Err(e) => {
1095 let msg = format!("{e}");
1096 assert!(
1097 msg.contains("totally-made-up") || msg.contains("Unknown provider"),
1098 "error message: {msg}"
1099 );
1100 },
1101 }
1102 }
1103
1104 #[test]
1105 fn normalize_cache_key_lowercases_provider_only() {
1106 assert_eq!(
1108 normalize_cache_key("Anthropic/Claude-X"),
1109 "anthropic/Claude-X"
1110 );
1111 assert_eq!(
1112 normalize_cache_key("anthropic/Claude-X"),
1113 "anthropic/Claude-X"
1114 );
1115 assert_eq!(normalize_cache_key("qwen3:30b"), "ollama/qwen3:30b");
1117 }
1118
1119 #[tokio::test]
1120 async fn resolve_is_single_flight_and_cached() {
1121 let f = ProviderFactory::new(Config::default());
1125 let (a, b) = tokio::join!(
1126 f.resolve("ollama/test-model"),
1127 f.resolve("Ollama/test-model"),
1128 );
1129 let a = a.expect("resolve a");
1130 let b = b.expect("resolve b");
1131 assert!(
1132 Arc::ptr_eq(&a, &b),
1133 "expected one cached provider for casing variants + concurrent resolve"
1134 );
1135 }
1136
1137 #[test]
1141 fn builtin_base_url_override_validated_and_resolved() {
1142 assert_eq!(
1144 resolve_overridable_base_url("anthropic", None, "https://api.anthropic.com/v1")
1145 .unwrap(),
1146 "https://api.anthropic.com/v1"
1147 );
1148 assert_eq!(
1150 resolve_overridable_base_url(
1151 "anthropic",
1152 Some("https://proxy.internal/v1".to_string()),
1153 "https://api.anthropic.com/v1",
1154 )
1155 .unwrap(),
1156 "https://proxy.internal/v1"
1157 );
1158 assert!(
1161 resolve_overridable_base_url(
1162 "anthropic",
1163 Some("http://attacker.example/v1".to_string()),
1164 "https://api.anthropic.com/v1",
1165 )
1166 .is_err()
1167 );
1168 assert!(
1170 resolve_overridable_base_url(
1171 "openai",
1172 Some("http://localhost:8080/v1".to_string()),
1173 "https://api.openai.com/v1",
1174 )
1175 .is_ok()
1176 );
1177 }
1178
1179 #[test]
1180 fn provider_host_extracts_host_or_unknown() {
1181 assert_eq!(
1182 provider_host("https://attacker.example/v1"),
1183 "attacker.example"
1184 );
1185 assert_eq!(provider_host("http://127.0.0.1:8080"), "127.0.0.1");
1186 assert_eq!(provider_host("not a url"), "<unknown>");
1187 }
1188
1189 #[test]
1190 fn override_host_warning_is_deduped() {
1191 let key = unique_env("MERMAID_FACTORY_WARN_KEY");
1195 assert!(should_warn_once(&key), "first warn for a key must fire");
1196 assert!(
1197 !should_warn_once(&key),
1198 "subsequent warns for the same key must be suppressed"
1199 );
1200 }
1201
1202 #[test]
1206 fn custom_profile_is_memoized_per_key() {
1207 use mermaid_domain::UserProviderConfig;
1208 let cfg = UserProviderConfig {
1209 base_url: Some("https://api.custom.test/v1".to_string()),
1210 api_key_env: Some("CUSTOM_KEY".to_string()),
1211 compat: Some("openai".to_string()),
1212 ..Default::default()
1213 };
1214 let a = user_profile_to_static("mermaid_test_customx", &cfg).unwrap();
1217 let b = user_profile_to_static("mermaid_test_customx", &cfg).unwrap();
1218 assert!(
1219 std::ptr::eq(a, b),
1220 "identical custom-provider inputs must reuse one leaked &'static profile"
1221 );
1222 assert_eq!(a.base_url, "https://api.custom.test/v1");
1223 assert_eq!(a.api_key_env, "CUSTOM_KEY");
1224
1225 let cfg2 = UserProviderConfig {
1227 base_url: Some("https://api.custom.test/v2".to_string()),
1228 ..cfg.clone()
1229 };
1230 let c = user_profile_to_static("mermaid_test_customx", &cfg2).unwrap();
1231 assert!(
1232 !std::ptr::eq(a, c),
1233 "a different base_url must leak a distinct profile"
1234 );
1235 }
1236}