1use std::sync::Arc;
9
10use tokio::sync::Mutex;
11
12use mermaid_domain::Config;
13use mermaid_model::models::config::BackendConfig;
14use mermaid_model::models::{ModelError, Result, lookup_provider};
15use mermaid_model::utils::{
16 resolve_api_key, resolve_provider_key, resolve_provider_key_with_fallback,
17};
18
19use super::model::{anthropic, gemini, meta};
20
21fn require_key(provider: &str, default_env: &str, override_env: Option<&str>) -> Result<String> {
25 resolve_provider_key(provider, default_env, override_env).ok_or_else(|| {
26 let env = override_env.unwrap_or(default_env);
27 ModelError::Authentication(format!(
28 "{provider} requires env var {env} (or `mermaid login {provider}`)"
29 ))
30 })
31}
32
33fn require_key_with_fallback(
34 provider: &str,
35 env_var: &str,
36 fallback_env_var: &str,
37) -> Result<String> {
38 resolve_provider_key_with_fallback(provider, env_var, fallback_env_var, None).ok_or_else(|| {
39 ModelError::Authentication(format!(
40 "{provider} requires env var {env_var} (or legacy {fallback_env_var}, or \
41 `mermaid login {provider}`)"
42 ))
43 })
44}
45
46fn key_env_source(default_env: &str, override_env: Option<&str>) -> Option<String> {
52 let env = override_env.unwrap_or(default_env);
56 resolve_api_key(env, None).map(|_| env.to_string())
57}
58
59fn key_env_source_with_fallback(default_env: &str, fallback_env: &str) -> Option<String> {
61 if resolve_api_key(default_env, None).is_some() {
62 return Some(default_env.to_string());
63 }
64 if resolve_api_key(fallback_env, None).is_some() {
65 return Some(fallback_env.to_string());
66 }
67 None
68}
69
70#[derive(Debug, Clone, PartialEq, Eq)]
80pub(crate) struct ProviderEndpoint {
81 pub(crate) base_url: String,
83 pub(crate) api_key: Option<String>,
86 pub(crate) key_env: Option<String>,
89}
90
91impl ProviderEndpoint {
92 fn require_key(&self, provider: &str) -> Result<String> {
98 self.api_key
99 .clone()
100 .ok_or_else(|| ModelError::Authentication(format!("{provider} requires an API key")))
101 }
102}
103
104fn is_known_provider(config: &Config, provider_lc: &str) -> bool {
107 matches!(provider_lc, "anthropic" | "gemini" | "meta")
108 || lookup_provider(provider_lc).is_some()
109 || config.providers.contains_key(provider_lc)
110}
111
112#[expect(
120 clippy::too_many_lines,
121 reason = "predates the lint; see .github/baselines/expect_budget.txt"
122)]
123pub(crate) fn resolve_provider_endpoint(
124 config: &Config,
125 provider: &str,
126) -> Result<ProviderEndpoint> {
127 let provider_lc = provider.to_lowercase();
128 let user_cfg = config.providers.get(&provider_lc);
129 let override_env = user_cfg.and_then(|c| c.api_key_env.as_deref());
130 let override_url = user_cfg.and_then(|c| c.base_url.clone());
131
132 let bespoke = match provider_lc.as_str() {
135 "anthropic" => Some((
136 anthropic::DEFAULT_API_KEY_ENV,
137 None,
138 anthropic::DEFAULT_BASE_URL,
139 )),
140 "gemini" => Some((
141 gemini::DEFAULT_API_KEY_ENV,
142 Some(gemini::LEGACY_API_KEY_ENV),
143 gemini::DEFAULT_BASE_URL,
144 )),
145 "meta" => Some((meta::DEFAULT_API_KEY_ENV, None, meta::DEFAULT_BASE_URL)),
146 _ => None,
147 };
148 if let Some((default_env, legacy_env, default_url)) = bespoke {
149 let base_url = resolve_overridable_base_url(&provider_lc, override_url, default_url)?;
150 let (api_key, key_env) = match legacy_env {
151 Some(legacy) if override_env.is_none() => (
154 require_key_with_fallback(&provider_lc, default_env, legacy)?,
155 key_env_source_with_fallback(default_env, legacy),
156 ),
157 _ => (
158 require_key(&provider_lc, default_env, override_env)?,
159 key_env_source(default_env, override_env),
160 ),
161 };
162 return Ok(ProviderEndpoint {
163 base_url,
164 api_key: Some(api_key),
165 key_env,
166 });
167 }
168
169 if provider_lc == "cloudflare" {
175 let profile = lookup_provider("cloudflare").expect("cloudflare is in the registry");
176 let api_key_env = override_env.unwrap_or(profile.api_key_env);
177 let base_url = match override_url {
178 Some(url) => {
181 validate_provider_base_url(&url)?;
182 warn_overridden_provider_host("cloudflare", &url);
183 url
184 },
185 None => match require_cloudflare_account_id() {
189 Ok(id) => cloudflare_base_url(&id),
190 Err(_)
191 if resolve_provider_key("cloudflare", profile.api_key_env, override_env)
192 .is_none() =>
193 {
194 return Err(ModelError::Authentication(format!(
195 "cloudflare requires env vars {api_key_env} and CLOUDFLARE_ACCOUNT_ID — \
196 create a token at https://dash.cloudflare.com/profile/api-tokens; the \
197 account id is on your Cloudflare dashboard (or set \
198 [providers.cloudflare].base_url)"
199 )));
200 },
201 Err(e) => return Err(e),
202 },
203 };
204 let api_key = resolve_optional_key(
205 &provider_lc,
206 profile.api_key_env,
207 override_env,
208 &base_url,
209 profile.key_hint,
210 )?;
211 let key_env = api_key
212 .is_some()
213 .then(|| key_env_source(profile.api_key_env, override_env))
214 .flatten();
215 return Ok(ProviderEndpoint {
216 base_url,
217 api_key,
218 key_env,
219 });
220 }
221
222 if let Some(profile) = lookup_provider(&provider_lc) {
224 let base_url = resolve_overridable_base_url(&provider_lc, override_url, profile.base_url)?;
225 let api_key = resolve_optional_key(
229 &provider_lc,
230 profile.api_key_env,
231 override_env,
232 &base_url,
233 profile.key_hint,
234 )?;
235 let key_env = api_key
236 .is_some()
237 .then(|| key_env_source(profile.api_key_env, override_env))
238 .flatten();
239 return Ok(ProviderEndpoint {
240 base_url,
241 api_key,
242 key_env,
243 });
244 }
245
246 if user_cfg.is_some() {
249 let base_url = override_url.ok_or_else(|| {
250 ModelError::InvalidRequest(format!(
251 "custom provider '{provider_lc}' requires base_url in config"
252 ))
253 })?;
254 let resolved = match override_env {
261 Some(env) => resolve_provider_key(&provider_lc, env, None),
262 None => mermaid_model::utils::default_store().get(&provider_lc),
263 };
264 let (api_key, key_env) = match resolved {
265 Some(key) => {
266 validate_provider_base_url(&base_url)?;
267 let env = override_env.and_then(|env| key_env_source(env, None));
268 (Some(key), env)
269 },
270 None if base_url_is_local(&base_url) => (None, None),
271 None => {
272 let reason = match override_env {
273 Some(env) => format!(
274 "requires env var {env} (or `mermaid login {provider_lc}`, or a \
275 loopback/LAN base_url)"
276 ),
277 None => "requires api_key_env, or a loopback/LAN base_url".to_string(),
278 };
279 return Err(ModelError::Authentication(format!(
280 "custom provider '{provider_lc}' {reason}"
281 )));
282 },
283 };
284 return Ok(ProviderEndpoint {
285 base_url,
286 api_key,
287 key_env,
288 });
289 }
290
291 Err(ModelError::InvalidRequest(format!(
292 "Unknown provider '{provider}'"
293 )))
294}
295
296fn resolve_optional_key(
301 provider: &str,
302 default_env: &str,
303 override_env: Option<&str>,
304 base_url: &str,
305 hint: Option<&str>,
306) -> Result<Option<String>> {
307 if let Some(key) = resolve_provider_key(provider, default_env, override_env) {
308 return Ok(Some(key));
309 }
310 if base_url_is_local(base_url) {
311 return Ok(None);
312 }
313 let env = override_env.unwrap_or(default_env);
314 let mut msg = format!("{provider} requires env var {env} (or `mermaid login {provider}`)");
315 if let Some(h) = hint {
316 msg.push_str(" — ");
317 msg.push_str(h);
318 }
319 Err(ModelError::Authentication(msg))
320}
321
322fn base_url_is_local(base_url: &str) -> bool {
325 reqwest::Url::parse(base_url)
326 .ok()
327 .and_then(|u| {
328 u.host_str()
329 .map(|h| mermaid_model::utils::classify_host(h).is_internal())
330 })
331 .unwrap_or(false)
332}
333
334fn merged_headers(
340 profile: &mermaid_model::models::ProviderProfile,
341 user_cfg: Option<&mermaid_domain::UserProviderConfig>,
342) -> std::collections::HashMap<String, String> {
343 let mut headers: std::collections::HashMap<String, String> = profile
344 .extra_headers
345 .iter()
346 .map(|(k, v)| ((*k).to_string(), (*v).to_string()))
347 .collect();
348 if let Some(cfg) = user_cfg {
349 for (k, v) in &cfg.extra_headers {
350 headers.insert(k.clone(), v.clone());
351 }
352 for (header, env_var) in &cfg.env_headers {
353 if let Ok(val) = std::env::var(env_var) {
354 headers.insert(header.clone(), val);
355 }
356 }
357 }
358 headers
359}
360
361use super::model::{
362 AnthropicProvider, GeminiProvider, MetaProvider, ModelProvider, OllamaProvider,
363 OpenAICompatProvider,
364};
365
366type ProviderCell = Arc<tokio::sync::OnceCell<Arc<dyn ModelProvider>>>;
370
371pub struct ProviderFactory {
375 config: Arc<Config>,
376 cache: Mutex<std::collections::HashMap<String, ProviderCell>>,
379}
380
381impl ProviderFactory {
382 #[must_use]
383 pub fn new(config: Config) -> Self {
384 Self {
385 config: Arc::new(config),
386 cache: Mutex::new(std::collections::HashMap::new()),
387 }
388 }
389
390 pub fn with_seeded_providers(
401 config: Config,
402 seeds: impl IntoIterator<Item = (String, Arc<dyn ModelProvider>)>,
403 ) -> Self {
404 let cache = seeds
405 .into_iter()
406 .map(|(model_id, provider)| {
407 let cell = tokio::sync::OnceCell::new_with(Some(provider));
408 (normalize_cache_key(&model_id), Arc::new(cell))
409 })
410 .collect();
411 Self {
412 config: Arc::new(config),
413 cache: Mutex::new(cache),
414 }
415 }
416
417 pub fn config(&self) -> &Config {
418 &self.config
419 }
420
421 pub async fn resolve(&self, model_id: &str) -> Result<Arc<dyn ModelProvider>> {
432 let key = normalize_cache_key(model_id);
433 let cell = {
437 let mut cache = self.cache.lock().await;
438 Arc::clone(
439 cache
440 .entry(key)
441 .or_insert_with(|| Arc::new(tokio::sync::OnceCell::new())),
442 )
443 };
444 let provider = cell
445 .get_or_try_init(|| async {
446 let p = build_provider(&self.config, model_id).await?;
447 Ok::<Arc<dyn ModelProvider>, ModelError>(Arc::from(p))
448 })
449 .await?;
450 Ok(Arc::clone(provider))
451 }
452}
453
454async fn build_provider(config: &Config, model_id: &str) -> Result<Box<dyn ModelProvider>> {
463 let (provider, model_name) = parse_model_id(model_id);
464 let provider_lc = provider.to_lowercase();
465
466 if provider_lc == "ollama" {
469 let backend = ollama_backend_config(config);
470 let p = OllamaProvider::with_app_config(
471 model_name,
472 Arc::new(backend),
473 Arc::new(config.clone()),
474 )
475 .await?;
476 return Ok(Box::new(p));
477 }
478
479 if !is_known_provider(config, &provider_lc) {
484 return Err(ModelError::InvalidRequest(format!(
485 "Unknown provider '{provider}' (model_id: {model_id})"
486 )));
487 }
488 let endpoint = resolve_provider_endpoint(config, &provider_lc)?;
489 let user_cfg = config.providers.get(&provider_lc);
490
491 if provider_lc == "anthropic" {
493 let p = AnthropicProvider::new(
494 endpoint.require_key(&provider_lc)?,
495 model_name.to_string(),
496 endpoint.base_url,
497 )?;
498 return Ok(Box::new(p));
499 }
500
501 if provider_lc == "gemini" {
503 let p = GeminiProvider::new(
504 endpoint.require_key(&provider_lc)?,
505 model_name.to_string(),
506 endpoint.base_url,
507 )?;
508 return Ok(Box::new(p));
509 }
510
511 if provider_lc == "meta" {
514 let api_key = endpoint.require_key(&provider_lc)?;
515 let mut extra_headers = std::collections::HashMap::new();
516 if let Some(cfg) = user_cfg {
517 extra_headers.extend(cfg.extra_headers.clone());
518 for (header, env_var) in &cfg.env_headers {
519 if let Ok(value) = std::env::var(env_var) {
520 extra_headers.insert(header.clone(), value);
521 }
522 }
523 }
524 let p = MetaProvider::new(
525 api_key,
526 model_name.to_string(),
527 endpoint.base_url,
528 extra_headers,
529 )?;
530 return Ok(Box::new(p));
531 }
532
533 let profile = match lookup_provider(&provider_lc) {
537 Some(profile) => profile,
538 None => user_cfg
541 .and_then(|cfg| user_profile_to_static(&provider_lc, cfg))
542 .ok_or_else(|| {
543 ModelError::InvalidRequest(format!(
544 "Unknown provider '{provider}' (model_id: {model_id})"
545 ))
546 })?,
547 };
548 let extra_headers = merged_headers(profile, user_cfg);
549 let p = OpenAICompatProvider::new(
550 profile,
551 endpoint.base_url,
552 endpoint.api_key,
553 model_name.to_string(),
554 extra_headers,
555 )?;
556 Ok(Box::new(p))
557}
558
559fn normalize_cache_key(model_id: &str) -> String {
563 let (provider, model) = parse_model_id(model_id);
564 format!("{}/{}", provider.to_lowercase(), model)
565}
566
567fn parse_model_id(model_id: &str) -> (String, &str) {
570 match model_id.split_once('/') {
571 Some((p, m)) => (p.to_string(), m),
572 None => ("ollama".to_string(), model_id),
573 }
574}
575
576pub(crate) fn ollama_backend_config(config: &Config) -> BackendConfig {
577 BackendConfig {
578 ollama_url: format!("{}:{}", config.ollama.host, config.ollama.port),
581 max_idle_per_host: 10,
582 timeout_secs: 10,
583 ollama_autostart: config.ollama.auto_start,
584 }
585}
586
587static PROFILE_CACHE: std::sync::LazyLock<
591 std::sync::Mutex<
592 std::collections::HashMap<String, &'static mermaid_model::models::ProviderProfile>,
593 >,
594> = std::sync::LazyLock::new(|| std::sync::Mutex::new(std::collections::HashMap::new()));
595
596fn user_profile_to_static(
610 name: &str,
611 user_cfg: &mermaid_domain::UserProviderConfig,
612) -> Option<&'static mermaid_model::models::ProviderProfile> {
613 use mermaid_model::models::{ProviderProfile, ReasoningExtraction, ReasoningStrategy};
614
615 let compat = user_cfg.compat.as_deref().unwrap_or("openai");
616 let base_url = user_cfg.base_url.clone().unwrap_or_default();
617 let api_key_env = user_cfg.api_key_env.clone().unwrap_or_default();
618
619 let cache_key = format!("{name}\u{0}{base_url}\u{0}{api_key_env}\u{0}{compat}");
622
623 let mut cache = PROFILE_CACHE
624 .lock()
625 .unwrap_or_else(|poisoned| poisoned.into_inner());
626 if let Some(profile) = cache.get(cache_key.as_str()) {
627 return Some(*profile);
630 }
631
632 let strategy = match compat {
633 "openai" => ReasoningStrategy::None,
634 "openai-effort" => ReasoningStrategy::Effort,
635 "openrouter" => ReasoningStrategy::OpenRouterShape,
636 _ => ReasoningStrategy::None,
637 };
638
639 let profile = Box::new(ProviderProfile {
640 name: Box::leak(name.to_string().into_boxed_str()),
641 base_url: Box::leak(base_url.into_boxed_str()),
642 api_key_env: Box::leak(api_key_env.into_boxed_str()),
643 key_hint: None,
644 extra_headers: &[],
645 reasoning_strategy: strategy,
646 reasoning_extraction: ReasoningExtraction::None,
647 max_tokens_param: mermaid_model::models::MaxTokensParam::MaxTokens,
648 disable_parallel_tool_calls_for: &[],
649 });
650 let leaked: &'static ProviderProfile = Box::leak(profile);
651 cache.insert(cache_key, leaked);
652 Some(leaked)
653}
654
655fn cloudflare_base_url(account_id: &str) -> String {
660 format!(
661 "https://api.cloudflare.com/client/v4/accounts/{}/ai/v1",
662 account_id.trim()
663 )
664}
665
666pub(crate) fn discovery_base_url(
674 profile: &mermaid_model::models::ProviderProfile,
675 override_url: Option<String>,
676) -> Option<String> {
677 if override_url.is_some() {
678 return override_url;
679 }
680 if profile.name == "cloudflare" {
681 return require_cloudflare_account_id()
682 .ok()
683 .map(|id| cloudflare_base_url(&id));
684 }
685 Some(profile.base_url.to_string())
686}
687
688fn require_cloudflare_account_id() -> Result<String> {
693 resolve_api_key("CLOUDFLARE_ACCOUNT_ID", None)
694 .map(|s| s.trim().to_string())
695 .filter(|s| !s.is_empty())
696 .ok_or_else(|| {
697 ModelError::Authentication(
698 "cloudflare requires env var CLOUDFLARE_ACCOUNT_ID (your Cloudflare account id) — \
699 find it on your Cloudflare dashboard, or set [providers.cloudflare].base_url"
700 .to_string(),
701 )
702 })
703}
704
705fn validate_provider_base_url(url: &str) -> Result<()> {
711 let parsed = reqwest::Url::parse(url).map_err(|e| {
712 ModelError::InvalidRequest(format!("invalid provider base_url '{url}': {e}"))
713 })?;
714 match parsed.scheme() {
715 "https" => Ok(()),
716 "http"
721 if mermaid_model::utils::classify_host(parsed.host_str().unwrap_or_default())
722 .is_loopback() =>
723 {
724 Ok(())
725 },
726 "http" => Err(ModelError::InvalidRequest(format!(
727 "provider base_url '{url}' uses http:// to a non-loopback host — refusing to send the \
728 API key in cleartext. Use https, or http://localhost for a local server."
729 ))),
730 other => Err(ModelError::InvalidRequest(format!(
731 "provider base_url '{url}' has unsupported scheme '{other}' (use http or https)"
732 ))),
733 }
734}
735
736fn resolve_overridable_base_url(
752 provider: &str,
753 override_url: Option<String>,
754 default_url: &str,
755) -> Result<String> {
756 match override_url {
757 Some(url) => {
758 validate_provider_base_url(&url)?;
759 warn_overridden_provider_host(provider, &url);
760 Ok(url)
761 },
762 None => Ok(default_url.to_string()),
763 }
764}
765
766static WARNED_OVERRIDE_HOSTS: std::sync::LazyLock<
770 std::sync::Mutex<std::collections::HashSet<String>>,
771> = std::sync::LazyLock::new(|| std::sync::Mutex::new(std::collections::HashSet::new()));
772
773fn warn_overridden_provider_host(provider: &str, base_url: &str) {
777 let host = provider_host(base_url);
778 if should_warn_once(&format!("{provider}@{host}")) {
779 tracing::warn!(
780 "built-in provider '{}' base_url overridden in config: the {} API key will be sent to \
781 host '{}' instead of the trusted default endpoint",
782 provider,
783 provider,
784 host
785 );
786 }
787}
788
789fn provider_host(base_url: &str) -> String {
792 reqwest::Url::parse(base_url)
793 .ok()
794 .and_then(|u| u.host_str().map(str::to_string))
795 .unwrap_or_else(|| "<unknown>".to_string())
796}
797
798fn should_warn_once(key: &str) -> bool {
801 let mut warned = WARNED_OVERRIDE_HOSTS
802 .lock()
803 .unwrap_or_else(|poisoned| poisoned.into_inner());
804 warned.insert(key.to_string())
805}
806
807#[cfg(test)]
808mod tests {
809 use super::*;
810
811 #[test]
812 fn base_url_is_local_classifies_hosts() {
813 assert!(base_url_is_local("http://127.0.0.1:8000/v1"));
814 assert!(base_url_is_local("http://localhost:1234/v1"));
815 assert!(base_url_is_local("http://192.168.1.5:8000/v1"));
816 assert!(!base_url_is_local("https://api.openai.com/v1"));
817 assert!(!base_url_is_local("not a url"));
818 }
819
820 #[test]
821 fn merged_headers_keeps_static_profile_headers_and_user_overrides() {
822 let profile = mermaid_model::models::lookup_provider("openrouter").unwrap();
823 let base = merged_headers(profile, None);
825 assert_eq!(
826 base.get("X-OpenRouter-Title").map(String::as_str),
827 Some("Mermaid")
828 );
829 assert!(base.contains_key("HTTP-Referer"));
830 let mut cfg = mermaid_domain::UserProviderConfig::default();
832 cfg.extra_headers.insert("X-Custom".into(), "v".into());
833 cfg.extra_headers
834 .insert("X-OpenRouter-Title".into(), "Override".into());
835 let merged = merged_headers(profile, Some(&cfg));
836 assert_eq!(merged.get("X-Custom").map(String::as_str), Some("v"));
837 assert_eq!(
838 merged.get("X-OpenRouter-Title").map(String::as_str),
839 Some("Override")
840 );
841 assert!(merged.contains_key("HTTP-Referer"));
842 }
843
844 #[test]
845 fn merged_headers_resolves_env_headers_and_skips_missing() {
846 let profile = mermaid_model::models::lookup_provider("openai").unwrap();
847 let mut cfg = mermaid_domain::UserProviderConfig::default();
848 cfg.env_headers
849 .insert("X-Gateway-Token".into(), "MERMAID_TEST_GW_TOKEN".into());
850 temp_env::with_var("MERMAID_TEST_GW_TOKEN", Some("secret123"), || {
851 let merged = merged_headers(profile, Some(&cfg));
852 assert_eq!(
853 merged.get("X-Gateway-Token").map(String::as_str),
854 Some("secret123")
855 );
856 });
857 temp_env::with_var("MERMAID_TEST_GW_TOKEN", None::<&str>, || {
858 assert!(!merged_headers(profile, Some(&cfg)).contains_key("X-Gateway-Token"));
859 });
860 }
861
862 #[test]
863 fn base_url_requires_https_for_remote_hosts() {
864 assert!(validate_provider_base_url("http://api.example.com/v1").is_err());
866 assert!(validate_provider_base_url("ftp://example.com").is_err());
867 assert!(validate_provider_base_url("https://api.example.com/v1").is_ok());
869 assert!(validate_provider_base_url("http://localhost:11434/v1").is_ok());
870 assert!(validate_provider_base_url("http://127.0.0.1:8000").is_ok());
871 assert!(validate_provider_base_url("http://[::1]:8000").is_ok());
872 assert!(validate_provider_base_url("http://192.168.1.5:8080").is_err());
875 assert!(validate_provider_base_url("http://169.254.169.254").is_err());
876 }
877
878 #[test]
879 fn cloudflare_base_url_synthesizes_account_scoped_endpoint() {
880 assert_eq!(
881 cloudflare_base_url("acct123"),
882 "https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1"
883 );
884 assert_eq!(
886 cloudflare_base_url(" acct123\n"),
887 "https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1"
888 );
889 }
890
891 #[test]
892 fn cloudflare_account_id_required_and_non_blank() {
893 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", None::<&str>, || {
895 let err = require_cloudflare_account_id().expect_err("must error when unset");
896 assert!(format!("{err}").contains("CLOUDFLARE_ACCOUNT_ID"));
897 });
898 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some(" "), || {
900 assert!(require_cloudflare_account_id().is_err());
901 });
902 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some(" acct123 "), || {
904 assert_eq!(require_cloudflare_account_id().unwrap(), "acct123");
905 });
906 }
907
908 #[test]
909 fn discovery_base_url_resolves_per_provider() {
910 let cf = lookup_provider("cloudflare").expect("cloudflare is in the registry");
911 let openai = lookup_provider("openai").expect("openai is in the registry");
912 assert_eq!(
914 discovery_base_url(cf, Some("https://gw.example/v1".to_string())),
915 Some("https://gw.example/v1".to_string())
916 );
917 assert_eq!(
919 discovery_base_url(openai, None),
920 Some(openai.base_url.to_string())
921 );
922 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some("acct123"), || {
924 assert_eq!(
925 discovery_base_url(cf, None),
926 Some("https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1".to_string())
927 );
928 });
929 temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", None::<&str>, || {
931 assert_eq!(discovery_base_url(cf, None), None);
932 });
933 }
934
935 #[tokio::test]
936 async fn cloudflare_missing_both_env_vars_is_one_combined_error() {
937 temp_env::async_with_vars(
938 [
939 ("CLOUDFLARE_ACCOUNT_ID", None::<&str>),
940 ("CLOUDFLARE_API_TOKEN", None),
941 ],
942 async {
943 let f = ProviderFactory::new(Config::default());
944 let err = match f.resolve("cloudflare/@cf/zai-org/glm-5.2").await {
945 Ok(_) => panic!("must fail with neither env var set"),
946 Err(e) => e,
947 };
948 let msg = format!("{err}");
949 assert!(
950 msg.contains("CLOUDFLARE_API_TOKEN") && msg.contains("CLOUDFLARE_ACCOUNT_ID"),
951 "one error must name both missing vars, got: {msg}"
952 );
953 },
954 )
955 .await;
956 }
957
958 use std::sync::atomic::{AtomicUsize, Ordering};
959
960 fn unique_env(prefix: &str) -> String {
961 static N: AtomicUsize = AtomicUsize::new(0);
962 format!(
963 "{}_{}_{}",
964 prefix,
965 std::process::id(),
966 N.fetch_add(1, Ordering::SeqCst)
967 )
968 }
969
970 #[test]
971 fn parse_bare_name_defaults_to_ollama() {
972 let (p, m) = parse_model_id("qwen3-coder:30b");
973 assert_eq!(p, "ollama");
974 assert_eq!(m, "qwen3-coder:30b");
975 }
976
977 #[test]
978 fn parse_prefixed() {
979 let (p, m) = parse_model_id("anthropic/claude-opus-4-7");
980 assert_eq!(p, "anthropic");
981 assert_eq!(m, "claude-opus-4-7");
982 }
983
984 #[tokio::test]
985 async fn meta_requires_its_documented_api_key_env() {
986 temp_env::async_with_vars(
987 [(
988 crate::providers::model::meta::DEFAULT_API_KEY_ENV,
989 None::<&str>,
990 )],
991 async {
992 let factory = ProviderFactory::new(Config::default());
993 let error = match factory.resolve("meta/muse-spark-1.1").await {
994 Ok(_) => panic!("Meta must require an API key"),
995 Err(error) => error,
996 };
997 assert!(
998 error
999 .to_string()
1000 .contains(crate::providers::model::meta::DEFAULT_API_KEY_ENV)
1001 );
1002 },
1003 )
1004 .await;
1005 }
1006
1007 #[tokio::test]
1008 async fn meta_routes_to_responses_provider_with_muse_capabilities() {
1009 temp_env::async_with_vars(
1010 [(
1011 crate::providers::model::meta::DEFAULT_API_KEY_ENV,
1012 Some("test-key"),
1013 )],
1014 async {
1015 let factory = ProviderFactory::new(Config::default());
1016 let provider = factory.resolve("meta/muse-spark-1.1").await.unwrap();
1017 let capabilities = provider.capabilities();
1018 assert!(capabilities.supports_tools);
1019 assert!(capabilities.supports_vision);
1020 assert!(capabilities.emits_provider_continuation);
1021 assert_eq!(
1022 capabilities.max_context_tokens,
1023 Some(mermaid_model::constants::META_MUSE_SPARK_CONTEXT_WINDOW)
1024 );
1025 assert_eq!(
1026 capabilities.max_output_tokens,
1027 Some(mermaid_model::constants::META_MUSE_SPARK_MAX_OUTPUT_TOKENS)
1028 );
1029 },
1030 )
1031 .await;
1032 }
1033
1034 #[test]
1035 fn gemini_key_resolution_accepts_legacy_fallback() {
1036 let primary = unique_env("MERMAID_FACTORY_GEMINI_PRIMARY");
1037 let legacy = unique_env("MERMAID_FACTORY_GEMINI_LEGACY");
1038 temp_env::with_vars(
1039 [(primary.as_str(), None), (legacy.as_str(), Some("legacy"))],
1040 || {
1041 let resolved = require_key_with_fallback("gemini", &primary, &legacy)
1042 .expect("legacy fallback should resolve");
1043 assert_eq!(resolved, "legacy");
1044 },
1045 );
1046 }
1047
1048 #[test]
1049 fn gemini_key_resolution_prefers_google_primary() {
1050 let primary = unique_env("MERMAID_FACTORY_GEMINI_PRIMARY2");
1051 let legacy = unique_env("MERMAID_FACTORY_GEMINI_LEGACY2");
1052 temp_env::with_vars(
1053 [
1054 (primary.as_str(), Some("google")),
1055 (legacy.as_str(), Some("legacy")),
1056 ],
1057 || {
1058 let resolved = require_key_with_fallback("gemini", &primary, &legacy)
1059 .expect("primary should resolve");
1060 assert_eq!(resolved, "google");
1061 },
1062 );
1063 }
1064
1065 #[tokio::test]
1066 async fn factory_reports_unknown_provider_clearly() {
1067 let cfg = Config::default();
1068 let f = ProviderFactory::new(cfg);
1069 match f.resolve("totally-made-up/model").await {
1070 Ok(_) => panic!("expected error"),
1071 Err(e) => {
1072 let msg = format!("{e}");
1073 assert!(
1074 msg.contains("totally-made-up") || msg.contains("Unknown provider"),
1075 "error message: {msg}"
1076 );
1077 },
1078 }
1079 }
1080
1081 #[test]
1082 fn normalize_cache_key_lowercases_provider_only() {
1083 assert_eq!(
1085 normalize_cache_key("Anthropic/Claude-X"),
1086 "anthropic/Claude-X"
1087 );
1088 assert_eq!(
1089 normalize_cache_key("anthropic/Claude-X"),
1090 "anthropic/Claude-X"
1091 );
1092 assert_eq!(normalize_cache_key("qwen3:30b"), "ollama/qwen3:30b");
1094 }
1095
1096 #[tokio::test]
1097 async fn resolve_is_single_flight_and_cached() {
1098 let f = ProviderFactory::new(Config::default());
1102 let (a, b) = tokio::join!(
1103 f.resolve("ollama/test-model"),
1104 f.resolve("Ollama/test-model"),
1105 );
1106 let a = a.expect("resolve a");
1107 let b = b.expect("resolve b");
1108 assert!(
1109 Arc::ptr_eq(&a, &b),
1110 "expected one cached provider for casing variants + concurrent resolve"
1111 );
1112 }
1113
1114 #[test]
1118 fn builtin_base_url_override_validated_and_resolved() {
1119 assert_eq!(
1121 resolve_overridable_base_url("anthropic", None, "https://api.anthropic.com/v1")
1122 .unwrap(),
1123 "https://api.anthropic.com/v1"
1124 );
1125 assert_eq!(
1127 resolve_overridable_base_url(
1128 "anthropic",
1129 Some("https://proxy.internal/v1".to_string()),
1130 "https://api.anthropic.com/v1",
1131 )
1132 .unwrap(),
1133 "https://proxy.internal/v1"
1134 );
1135 assert!(
1138 resolve_overridable_base_url(
1139 "anthropic",
1140 Some("http://attacker.example/v1".to_string()),
1141 "https://api.anthropic.com/v1",
1142 )
1143 .is_err()
1144 );
1145 assert!(
1147 resolve_overridable_base_url(
1148 "openai",
1149 Some("http://localhost:8080/v1".to_string()),
1150 "https://api.openai.com/v1",
1151 )
1152 .is_ok()
1153 );
1154 }
1155
1156 #[test]
1157 fn provider_host_extracts_host_or_unknown() {
1158 assert_eq!(
1159 provider_host("https://attacker.example/v1"),
1160 "attacker.example"
1161 );
1162 assert_eq!(provider_host("http://127.0.0.1:8080"), "127.0.0.1");
1163 assert_eq!(provider_host("not a url"), "<unknown>");
1164 }
1165
1166 #[test]
1167 fn override_host_warning_is_deduped() {
1168 let key = unique_env("MERMAID_FACTORY_WARN_KEY");
1172 assert!(should_warn_once(&key), "first warn for a key must fire");
1173 assert!(
1174 !should_warn_once(&key),
1175 "subsequent warns for the same key must be suppressed"
1176 );
1177 }
1178
1179 #[test]
1183 fn custom_profile_is_memoized_per_key() {
1184 use mermaid_domain::UserProviderConfig;
1185 let cfg = UserProviderConfig {
1186 base_url: Some("https://api.custom.test/v1".to_string()),
1187 api_key_env: Some("CUSTOM_KEY".to_string()),
1188 compat: Some("openai".to_string()),
1189 ..Default::default()
1190 };
1191 let a = user_profile_to_static("mermaid_test_customx", &cfg).unwrap();
1194 let b = user_profile_to_static("mermaid_test_customx", &cfg).unwrap();
1195 assert!(
1196 std::ptr::eq(a, b),
1197 "identical custom-provider inputs must reuse one leaked &'static profile"
1198 );
1199 assert_eq!(a.base_url, "https://api.custom.test/v1");
1200 assert_eq!(a.api_key_env, "CUSTOM_KEY");
1201
1202 let cfg2 = UserProviderConfig {
1204 base_url: Some("https://api.custom.test/v2".to_string()),
1205 ..cfg.clone()
1206 };
1207 let c = user_profile_to_static("mermaid_test_customx", &cfg2).unwrap();
1208 assert!(
1209 !std::ptr::eq(a, c),
1210 "a different base_url must leak a distinct profile"
1211 );
1212 }
1213}