use std::sync::Arc;
use tokio::sync::Mutex;
use crate::app::Config;
use crate::models::config::BackendConfig;
use crate::models::{ModelError, Result, lookup_provider};
use crate::utils::{resolve_api_key, resolve_provider_key, resolve_provider_key_with_fallback};
const GEMINI_API_KEY_ENV: &str = "GOOGLE_API_KEY";
const GEMINI_LEGACY_API_KEY_ENV: &str = "GEMINI_API_KEY";
fn require_key(provider: &str, default_env: &str, override_env: Option<&str>) -> Result<String> {
resolve_provider_key(provider, default_env, override_env).ok_or_else(|| {
let env = override_env.unwrap_or(default_env);
ModelError::Authentication(format!(
"{provider} requires env var {env} (or `mermaid login {provider}`)"
))
})
}
fn require_key_with_fallback(
provider: &str,
env_var: &str,
fallback_env_var: &str,
) -> Result<String> {
resolve_provider_key_with_fallback(provider, env_var, fallback_env_var, None).ok_or_else(|| {
ModelError::Authentication(format!(
"{provider} requires env var {env_var} (or legacy {fallback_env_var}, or \
`mermaid login {provider}`)"
))
})
}
fn resolve_optional_key(
provider: &str,
default_env: &str,
override_env: Option<&str>,
base_url: &str,
hint: Option<&str>,
) -> Result<Option<String>> {
if let Some(key) = resolve_provider_key(provider, default_env, override_env) {
return Ok(Some(key));
}
if base_url_is_local(base_url) {
return Ok(None);
}
let env = override_env.unwrap_or(default_env);
let mut msg = format!("{provider} requires env var {env} (or `mermaid login {provider}`)");
if let Some(h) = hint {
msg.push_str(" — ");
msg.push_str(h);
}
Err(ModelError::Authentication(msg))
}
fn base_url_is_local(base_url: &str) -> bool {
reqwest::Url::parse(base_url)
.ok()
.and_then(|u| {
u.host_str()
.map(|h| crate::utils::classify_host(h).is_internal())
})
.unwrap_or(false)
}
fn merged_headers(
profile: &crate::models::ProviderProfile,
user_cfg: Option<&crate::app::UserProviderConfig>,
) -> std::collections::HashMap<String, String> {
let mut headers: std::collections::HashMap<String, String> = profile
.extra_headers
.iter()
.map(|(k, v)| ((*k).to_string(), (*v).to_string()))
.collect();
if let Some(cfg) = user_cfg {
for (k, v) in &cfg.extra_headers {
headers.insert(k.clone(), v.clone());
}
for (header, env_var) in &cfg.env_headers {
if let Ok(val) = std::env::var(env_var) {
headers.insert(header.clone(), val);
}
}
}
headers
}
use super::model::{
AnthropicProvider, GeminiProvider, MetaProvider, ModelProvider, OllamaProvider,
OpenAICompatProvider,
};
type ProviderCell = Arc<tokio::sync::OnceCell<Arc<dyn ModelProvider>>>;
pub struct ProviderFactory {
config: Arc<Config>,
cache: Mutex<std::collections::HashMap<String, ProviderCell>>,
}
impl ProviderFactory {
pub fn new(config: Config) -> Self {
Self {
config: Arc::new(config),
cache: Mutex::new(std::collections::HashMap::new()),
}
}
pub fn config(&self) -> &Config {
&self.config
}
pub async fn resolve(&self, model_id: &str) -> Result<Arc<dyn ModelProvider>> {
let key = normalize_cache_key(model_id);
let cell = {
let mut cache = self.cache.lock().await;
Arc::clone(
cache
.entry(key)
.or_insert_with(|| Arc::new(tokio::sync::OnceCell::new())),
)
};
let provider = cell
.get_or_try_init(|| async {
let p = build_provider(&self.config, model_id).await?;
Ok::<Arc<dyn ModelProvider>, ModelError>(Arc::from(p))
})
.await?;
Ok(Arc::clone(provider))
}
}
async fn build_provider(config: &Config, model_id: &str) -> Result<Box<dyn ModelProvider>> {
let (provider, model_name) = parse_model_id(model_id);
let provider_lc = provider.to_lowercase();
if provider_lc == "ollama" {
let backend = ollama_backend_config(config);
let p = OllamaProvider::with_app_config(
model_name,
Arc::new(backend),
Arc::new(config.clone()),
)
.await?;
return Ok(Box::new(p));
}
if provider_lc == "anthropic" {
let user_cfg = config.providers.get("anthropic");
let base_url = resolve_overridable_base_url(
"anthropic",
user_cfg.and_then(|c| c.base_url.clone()),
"https://api.anthropic.com/v1",
)?;
let api_key = require_key(
"anthropic",
"ANTHROPIC_API_KEY",
user_cfg.and_then(|c| c.api_key_env.as_deref()),
)?;
let p = AnthropicProvider::new(api_key, model_name.to_string(), base_url)?;
return Ok(Box::new(p));
}
if provider_lc == "gemini" {
let user_cfg = config.providers.get("gemini");
let base_url = resolve_overridable_base_url(
"gemini",
user_cfg.and_then(|c| c.base_url.clone()),
"https://generativelanguage.googleapis.com/v1beta",
)?;
let api_key = match user_cfg.and_then(|c| c.api_key_env.as_deref()) {
Some(api_key_env) => require_key("gemini", GEMINI_API_KEY_ENV, Some(api_key_env))?,
None => {
require_key_with_fallback("gemini", GEMINI_API_KEY_ENV, GEMINI_LEGACY_API_KEY_ENV)?
},
};
let p = GeminiProvider::new(api_key, model_name.to_string(), base_url)?;
return Ok(Box::new(p));
}
if provider_lc == "meta" {
let user_cfg = config.providers.get("meta");
let base_url = resolve_overridable_base_url(
"meta",
user_cfg.and_then(|cfg| cfg.base_url.clone()),
super::model::meta::DEFAULT_BASE_URL,
)?;
let api_key = require_key(
"meta",
super::model::meta::DEFAULT_API_KEY_ENV,
user_cfg.and_then(|cfg| cfg.api_key_env.as_deref()),
)?;
let mut extra_headers = std::collections::HashMap::new();
if let Some(cfg) = user_cfg {
extra_headers.extend(cfg.extra_headers.clone());
for (header, env_var) in &cfg.env_headers {
if let Ok(value) = std::env::var(env_var) {
extra_headers.insert(header.clone(), value);
}
}
}
let p = MetaProvider::new(api_key, model_name.to_string(), base_url, extra_headers)?;
return Ok(Box::new(p));
}
if provider_lc == "cloudflare" {
let user_cfg = config.providers.get("cloudflare");
let profile = lookup_provider("cloudflare").expect("cloudflare is in the registry");
let override_env = user_cfg.and_then(|c| c.api_key_env.as_deref());
let api_key_env = override_env.unwrap_or(profile.api_key_env);
let base_url = match user_cfg.and_then(|c| c.base_url.clone()) {
Some(url) => {
validate_provider_base_url(&url)?;
warn_overridden_provider_host("cloudflare", &url);
url
},
None => match require_cloudflare_account_id() {
Ok(id) => cloudflare_base_url(&id),
Err(_)
if resolve_provider_key("cloudflare", profile.api_key_env, override_env)
.is_none() =>
{
return Err(ModelError::Authentication(format!(
"cloudflare requires env vars {api_key_env} and CLOUDFLARE_ACCOUNT_ID — \
create a token at https://dash.cloudflare.com/profile/api-tokens; the \
account id is on your Cloudflare dashboard (or set \
[providers.cloudflare].base_url)"
)));
},
Err(e) => return Err(e),
},
};
let api_key = resolve_optional_key(
&provider_lc,
profile.api_key_env,
override_env,
&base_url,
profile.key_hint,
)?;
let extra_headers = merged_headers(profile, user_cfg);
let p = OpenAICompatProvider::new(
profile,
base_url,
api_key,
model_name.to_string(),
extra_headers,
)?;
return Ok(Box::new(p));
}
if let Some(profile) = lookup_provider(&provider_lc) {
let user_cfg = config.providers.get(&provider_lc);
let base_url = resolve_overridable_base_url(
&provider_lc,
user_cfg.and_then(|c| c.base_url.clone()),
profile.base_url,
)?;
let api_key = resolve_optional_key(
&provider_lc,
profile.api_key_env,
user_cfg.and_then(|c| c.api_key_env.as_deref()),
&base_url,
profile.key_hint,
)?;
let extra_headers = merged_headers(profile, user_cfg);
let p = OpenAICompatProvider::new(
profile,
base_url,
api_key,
model_name.to_string(),
extra_headers,
)?;
return Ok(Box::new(p));
}
if let Some(user_cfg) = config.providers.get(&provider_lc)
&& let Some(profile) = user_profile_to_static(&provider_lc, user_cfg)
{
let base_url = user_cfg.base_url.clone().ok_or_else(|| {
ModelError::InvalidRequest(format!(
"custom provider '{}' requires base_url in config",
provider_lc
))
})?;
let api_key_env = user_cfg.api_key_env.as_deref();
let resolved = match api_key_env {
Some(env) => resolve_provider_key(&provider_lc, env, None),
None => crate::utils::default_store().get(&provider_lc),
};
let api_key = match resolved {
Some(key) => {
validate_provider_base_url(&base_url)?;
Some(key)
},
None if base_url_is_local(&base_url) => None,
None => {
let reason = match api_key_env {
Some(env) => format!(
"requires env var {env} (or `mermaid login {provider_lc}`, or a \
loopback/LAN base_url)"
),
None => "requires api_key_env, or a loopback/LAN base_url".to_string(),
};
return Err(ModelError::Authentication(format!(
"custom provider '{provider_lc}' {reason}"
)));
},
};
let extra_headers = merged_headers(profile, Some(user_cfg));
let p = OpenAICompatProvider::new(
profile,
base_url,
api_key,
model_name.to_string(),
extra_headers,
)?;
return Ok(Box::new(p));
}
Err(ModelError::InvalidRequest(format!(
"Unknown provider '{}' (model_id: {})",
provider, model_id
)))
}
fn normalize_cache_key(model_id: &str) -> String {
let (provider, model) = parse_model_id(model_id);
format!("{}/{}", provider.to_lowercase(), model)
}
fn parse_model_id(model_id: &str) -> (String, &str) {
match model_id.split_once('/') {
Some((p, m)) => (p.to_string(), m),
None => ("ollama".to_string(), model_id),
}
}
pub(crate) fn ollama_backend_config(config: &Config) -> BackendConfig {
BackendConfig {
ollama_url: format!("{}:{}", config.ollama.host, config.ollama.port),
max_idle_per_host: 10,
timeout_secs: 10,
ollama_autostart: config.ollama.auto_start,
}
}
static PROFILE_CACHE: std::sync::LazyLock<
std::sync::Mutex<std::collections::HashMap<String, &'static crate::models::ProviderProfile>>,
> = std::sync::LazyLock::new(|| std::sync::Mutex::new(std::collections::HashMap::new()));
fn user_profile_to_static(
name: &str,
user_cfg: &crate::app::UserProviderConfig,
) -> Option<&'static crate::models::ProviderProfile> {
use crate::models::{ProviderProfile, ReasoningExtraction, ReasoningStrategy};
let compat = user_cfg.compat.as_deref().unwrap_or("openai");
let base_url = user_cfg.base_url.clone().unwrap_or_default();
let api_key_env = user_cfg.api_key_env.clone().unwrap_or_default();
let cache_key = format!("{name}\u{0}{base_url}\u{0}{api_key_env}\u{0}{compat}");
let mut cache = PROFILE_CACHE
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(profile) = cache.get(cache_key.as_str()) {
return Some(*profile);
}
let strategy = match compat {
"openai" => ReasoningStrategy::None,
"openai-effort" => ReasoningStrategy::Effort,
"openrouter" => ReasoningStrategy::OpenRouterShape,
_ => ReasoningStrategy::None,
};
let profile = Box::new(ProviderProfile {
name: Box::leak(name.to_string().into_boxed_str()),
base_url: Box::leak(base_url.into_boxed_str()),
api_key_env: Box::leak(api_key_env.into_boxed_str()),
key_hint: None,
extra_headers: &[],
reasoning_strategy: strategy,
reasoning_extraction: ReasoningExtraction::None,
max_tokens_param: crate::models::MaxTokensParam::MaxTokens,
disable_parallel_tool_calls_for: &[],
});
let leaked: &'static ProviderProfile = Box::leak(profile);
cache.insert(cache_key, leaked);
Some(leaked)
}
fn cloudflare_base_url(account_id: &str) -> String {
format!(
"https://api.cloudflare.com/client/v4/accounts/{}/ai/v1",
account_id.trim()
)
}
pub(crate) fn discovery_base_url(
profile: &crate::models::ProviderProfile,
override_url: Option<String>,
) -> Option<String> {
if override_url.is_some() {
return override_url;
}
if profile.name == "cloudflare" {
return require_cloudflare_account_id()
.ok()
.map(|id| cloudflare_base_url(&id));
}
Some(profile.base_url.to_string())
}
fn require_cloudflare_account_id() -> Result<String> {
resolve_api_key("CLOUDFLARE_ACCOUNT_ID", None)
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.ok_or_else(|| {
ModelError::Authentication(
"cloudflare requires env var CLOUDFLARE_ACCOUNT_ID (your Cloudflare account id) — \
find it on your Cloudflare dashboard, or set [providers.cloudflare].base_url"
.to_string(),
)
})
}
fn validate_provider_base_url(url: &str) -> Result<()> {
let parsed = reqwest::Url::parse(url).map_err(|e| {
ModelError::InvalidRequest(format!("invalid provider base_url '{url}': {e}"))
})?;
match parsed.scheme() {
"https" => Ok(()),
"http"
if crate::utils::classify_host(parsed.host_str().unwrap_or_default()).is_loopback() =>
{
Ok(())
},
"http" => Err(ModelError::InvalidRequest(format!(
"provider base_url '{url}' uses http:// to a non-loopback host — refusing to send the \
API key in cleartext. Use https, or http://localhost for a local server."
))),
other => Err(ModelError::InvalidRequest(format!(
"provider base_url '{url}' has unsupported scheme '{other}' (use http or https)"
))),
}
}
fn resolve_overridable_base_url(
provider: &str,
override_url: Option<String>,
default_url: &str,
) -> Result<String> {
match override_url {
Some(url) => {
validate_provider_base_url(&url)?;
warn_overridden_provider_host(provider, &url);
Ok(url)
},
None => Ok(default_url.to_string()),
}
}
static WARNED_OVERRIDE_HOSTS: std::sync::LazyLock<
std::sync::Mutex<std::collections::HashSet<String>>,
> = std::sync::LazyLock::new(|| std::sync::Mutex::new(std::collections::HashSet::new()));
fn warn_overridden_provider_host(provider: &str, base_url: &str) {
let host = provider_host(base_url);
if should_warn_once(&format!("{provider}@{host}")) {
tracing::warn!(
"built-in provider '{}' base_url overridden in config: the {} API key will be sent to \
host '{}' instead of the trusted default endpoint",
provider,
provider,
host
);
}
}
fn provider_host(base_url: &str) -> String {
reqwest::Url::parse(base_url)
.ok()
.and_then(|u| u.host_str().map(str::to_string))
.unwrap_or_else(|| "<unknown>".to_string())
}
fn should_warn_once(key: &str) -> bool {
let mut warned = WARNED_OVERRIDE_HOSTS
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
warned.insert(key.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn base_url_is_local_classifies_hosts() {
assert!(base_url_is_local("http://127.0.0.1:8000/v1"));
assert!(base_url_is_local("http://localhost:1234/v1"));
assert!(base_url_is_local("http://192.168.1.5:8000/v1"));
assert!(!base_url_is_local("https://api.openai.com/v1"));
assert!(!base_url_is_local("not a url"));
}
#[test]
fn merged_headers_keeps_static_profile_headers_and_user_overrides() {
let profile = crate::models::lookup_provider("openrouter").unwrap();
let base = merged_headers(profile, None);
assert_eq!(
base.get("X-OpenRouter-Title").map(String::as_str),
Some("Mermaid")
);
assert!(base.contains_key("HTTP-Referer"));
let mut cfg = crate::app::UserProviderConfig::default();
cfg.extra_headers.insert("X-Custom".into(), "v".into());
cfg.extra_headers
.insert("X-OpenRouter-Title".into(), "Override".into());
let merged = merged_headers(profile, Some(&cfg));
assert_eq!(merged.get("X-Custom").map(String::as_str), Some("v"));
assert_eq!(
merged.get("X-OpenRouter-Title").map(String::as_str),
Some("Override")
);
assert!(merged.contains_key("HTTP-Referer"));
}
#[test]
fn merged_headers_resolves_env_headers_and_skips_missing() {
let profile = crate::models::lookup_provider("openai").unwrap();
let mut cfg = crate::app::UserProviderConfig::default();
cfg.env_headers
.insert("X-Gateway-Token".into(), "MERMAID_TEST_GW_TOKEN".into());
temp_env::with_var("MERMAID_TEST_GW_TOKEN", Some("secret123"), || {
let merged = merged_headers(profile, Some(&cfg));
assert_eq!(
merged.get("X-Gateway-Token").map(String::as_str),
Some("secret123")
);
});
temp_env::with_var("MERMAID_TEST_GW_TOKEN", None::<&str>, || {
assert!(!merged_headers(profile, Some(&cfg)).contains_key("X-Gateway-Token"));
});
}
#[test]
fn base_url_requires_https_for_remote_hosts() {
assert!(validate_provider_base_url("http://api.example.com/v1").is_err());
assert!(validate_provider_base_url("ftp://example.com").is_err());
assert!(validate_provider_base_url("https://api.example.com/v1").is_ok());
assert!(validate_provider_base_url("http://localhost:11434/v1").is_ok());
assert!(validate_provider_base_url("http://127.0.0.1:8000").is_ok());
assert!(validate_provider_base_url("http://[::1]:8000").is_ok());
assert!(validate_provider_base_url("http://192.168.1.5:8080").is_err());
assert!(validate_provider_base_url("http://169.254.169.254").is_err());
}
#[test]
fn cloudflare_base_url_synthesizes_account_scoped_endpoint() {
assert_eq!(
cloudflare_base_url("acct123"),
"https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1"
);
assert_eq!(
cloudflare_base_url(" acct123\n"),
"https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1"
);
}
#[test]
fn cloudflare_account_id_required_and_non_blank() {
temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", None::<&str>, || {
let err = require_cloudflare_account_id().expect_err("must error when unset");
assert!(format!("{err}").contains("CLOUDFLARE_ACCOUNT_ID"));
});
temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some(" "), || {
assert!(require_cloudflare_account_id().is_err());
});
temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some(" acct123 "), || {
assert_eq!(require_cloudflare_account_id().unwrap(), "acct123");
});
}
#[test]
fn discovery_base_url_resolves_per_provider() {
let cf = lookup_provider("cloudflare").expect("cloudflare is in the registry");
let openai = lookup_provider("openai").expect("openai is in the registry");
assert_eq!(
discovery_base_url(cf, Some("https://gw.example/v1".to_string())),
Some("https://gw.example/v1".to_string())
);
assert_eq!(
discovery_base_url(openai, None),
Some(openai.base_url.to_string())
);
temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", Some("acct123"), || {
assert_eq!(
discovery_base_url(cf, None),
Some("https://api.cloudflare.com/client/v4/accounts/acct123/ai/v1".to_string())
);
});
temp_env::with_var("CLOUDFLARE_ACCOUNT_ID", None::<&str>, || {
assert_eq!(discovery_base_url(cf, None), None);
});
}
#[tokio::test]
async fn cloudflare_missing_both_env_vars_is_one_combined_error() {
temp_env::async_with_vars(
[
("CLOUDFLARE_ACCOUNT_ID", None::<&str>),
("CLOUDFLARE_API_TOKEN", None),
],
async {
let f = ProviderFactory::new(Config::default());
let err = match f.resolve("cloudflare/@cf/zai-org/glm-5.2").await {
Ok(_) => panic!("must fail with neither env var set"),
Err(e) => e,
};
let msg = format!("{err}");
assert!(
msg.contains("CLOUDFLARE_API_TOKEN") && msg.contains("CLOUDFLARE_ACCOUNT_ID"),
"one error must name both missing vars, got: {msg}"
);
},
)
.await;
}
use std::sync::atomic::{AtomicUsize, Ordering};
fn unique_env(prefix: &str) -> String {
static N: AtomicUsize = AtomicUsize::new(0);
format!(
"{}_{}_{}",
prefix,
std::process::id(),
N.fetch_add(1, Ordering::SeqCst)
)
}
#[test]
fn parse_bare_name_defaults_to_ollama() {
let (p, m) = parse_model_id("qwen3-coder:30b");
assert_eq!(p, "ollama");
assert_eq!(m, "qwen3-coder:30b");
}
#[test]
fn parse_prefixed() {
let (p, m) = parse_model_id("anthropic/claude-opus-4-7");
assert_eq!(p, "anthropic");
assert_eq!(m, "claude-opus-4-7");
}
#[tokio::test]
async fn meta_requires_its_documented_api_key_env() {
temp_env::async_with_vars(
[(
crate::providers::model::meta::DEFAULT_API_KEY_ENV,
None::<&str>,
)],
async {
let factory = ProviderFactory::new(Config::default());
let error = match factory.resolve("meta/muse-spark-1.1").await {
Ok(_) => panic!("Meta must require an API key"),
Err(error) => error,
};
assert!(
error
.to_string()
.contains(crate::providers::model::meta::DEFAULT_API_KEY_ENV)
);
},
)
.await;
}
#[tokio::test]
async fn meta_routes_to_responses_provider_with_muse_capabilities() {
temp_env::async_with_vars(
[(
crate::providers::model::meta::DEFAULT_API_KEY_ENV,
Some("test-key"),
)],
async {
let factory = ProviderFactory::new(Config::default());
let provider = factory.resolve("meta/muse-spark-1.1").await.unwrap();
let capabilities = provider.capabilities();
assert!(capabilities.supports_tools);
assert!(capabilities.supports_vision);
assert!(capabilities.emits_provider_continuation);
assert_eq!(
capabilities.max_context_tokens,
Some(crate::constants::META_MUSE_SPARK_CONTEXT_WINDOW)
);
assert_eq!(
capabilities.max_output_tokens,
Some(crate::constants::META_MUSE_SPARK_MAX_OUTPUT_TOKENS)
);
},
)
.await;
}
#[test]
fn gemini_key_resolution_accepts_legacy_fallback() {
let primary = unique_env("MERMAID_FACTORY_GEMINI_PRIMARY");
let legacy = unique_env("MERMAID_FACTORY_GEMINI_LEGACY");
temp_env::with_vars(
[(primary.as_str(), None), (legacy.as_str(), Some("legacy"))],
|| {
let resolved = require_key_with_fallback("gemini", &primary, &legacy)
.expect("legacy fallback should resolve");
assert_eq!(resolved, "legacy");
},
);
}
#[test]
fn gemini_key_resolution_prefers_google_primary() {
let primary = unique_env("MERMAID_FACTORY_GEMINI_PRIMARY2");
let legacy = unique_env("MERMAID_FACTORY_GEMINI_LEGACY2");
temp_env::with_vars(
[
(primary.as_str(), Some("google")),
(legacy.as_str(), Some("legacy")),
],
|| {
let resolved = require_key_with_fallback("gemini", &primary, &legacy)
.expect("primary should resolve");
assert_eq!(resolved, "google");
},
);
}
#[tokio::test]
async fn factory_reports_unknown_provider_clearly() {
let cfg = Config::default();
let f = ProviderFactory::new(cfg);
match f.resolve("totally-made-up/model").await {
Ok(_) => panic!("expected error"),
Err(e) => {
let msg = format!("{}", e);
assert!(
msg.contains("totally-made-up") || msg.contains("Unknown provider"),
"error message: {}",
msg
);
},
}
}
#[test]
fn normalize_cache_key_lowercases_provider_only() {
assert_eq!(
normalize_cache_key("Anthropic/Claude-X"),
"anthropic/Claude-X"
);
assert_eq!(
normalize_cache_key("anthropic/Claude-X"),
"anthropic/Claude-X"
);
assert_eq!(normalize_cache_key("qwen3:30b"), "ollama/qwen3:30b");
}
#[tokio::test]
async fn resolve_is_single_flight_and_cached() {
let f = ProviderFactory::new(Config::default());
let (a, b) = tokio::join!(
f.resolve("ollama/test-model"),
f.resolve("Ollama/test-model"),
);
let a = a.expect("resolve a");
let b = b.expect("resolve b");
assert!(
Arc::ptr_eq(&a, &b),
"expected one cached provider for casing variants + concurrent resolve"
);
}
#[test]
fn builtin_base_url_override_validated_and_resolved() {
assert_eq!(
resolve_overridable_base_url("anthropic", None, "https://api.anthropic.com/v1")
.unwrap(),
"https://api.anthropic.com/v1"
);
assert_eq!(
resolve_overridable_base_url(
"anthropic",
Some("https://proxy.internal/v1".to_string()),
"https://api.anthropic.com/v1",
)
.unwrap(),
"https://proxy.internal/v1"
);
assert!(
resolve_overridable_base_url(
"anthropic",
Some("http://attacker.example/v1".to_string()),
"https://api.anthropic.com/v1",
)
.is_err()
);
assert!(
resolve_overridable_base_url(
"openai",
Some("http://localhost:8080/v1".to_string()),
"https://api.openai.com/v1",
)
.is_ok()
);
}
#[test]
fn provider_host_extracts_host_or_unknown() {
assert_eq!(
provider_host("https://attacker.example/v1"),
"attacker.example"
);
assert_eq!(provider_host("http://127.0.0.1:8080"), "127.0.0.1");
assert_eq!(provider_host("not a url"), "<unknown>");
}
#[test]
fn override_host_warning_is_deduped() {
let key = unique_env("MERMAID_FACTORY_WARN_KEY");
assert!(should_warn_once(&key), "first warn for a key must fire");
assert!(
!should_warn_once(&key),
"subsequent warns for the same key must be suppressed"
);
}
#[test]
fn custom_profile_is_memoized_per_key() {
use crate::app::UserProviderConfig;
let cfg = UserProviderConfig {
base_url: Some("https://api.custom.test/v1".to_string()),
api_key_env: Some("CUSTOM_KEY".to_string()),
compat: Some("openai".to_string()),
..Default::default()
};
let a = user_profile_to_static("mermaid_test_customx", &cfg).unwrap();
let b = user_profile_to_static("mermaid_test_customx", &cfg).unwrap();
assert!(
std::ptr::eq(a, b),
"identical custom-provider inputs must reuse one leaked &'static profile"
);
assert_eq!(a.base_url, "https://api.custom.test/v1");
assert_eq!(a.api_key_env, "CUSTOM_KEY");
let cfg2 = UserProviderConfig {
base_url: Some("https://api.custom.test/v2".to_string()),
..cfg.clone()
};
let c = user_profile_to_static("mermaid_test_customx", &cfg2).unwrap();
assert!(
!std::ptr::eq(a, c),
"a different base_url must leak a distinct profile"
);
}
}