use mentra::BuiltinProvider;
use thiserror::Error;
const CANDIDATES: &[(BuiltinProvider, &str)] = &[
(BuiltinProvider::Anthropic, "ANTHROPIC_API_KEY"),
(BuiltinProvider::OpenAI, "OPENAI_API_KEY"),
(BuiltinProvider::Gemini, "GEMINI_API_KEY"),
(BuiltinProvider::OpenRouter, "OPENROUTER_API_KEY"),
];
const BASE_URL_VARS: &[&str] = &["BASIS_BASE_URL", "OPENAI_BASE_URL"];
const COMPATIBLE_KEY_VARS: &[&str] = &["BASIS_API_KEY", "OPENAI_API_KEY"];
#[derive(Clone)]
pub struct ProviderChoice {
pub provider: BuiltinProvider,
pub api_key: String,
pub source_var: Option<&'static str>,
pub base_url: Option<String>,
}
impl std::fmt::Debug for ProviderChoice {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ProviderChoice")
.field("provider", &self.provider)
.field("api_key", &"<redacted>")
.field("source_var", &self.source_var)
.field("base_url", &self.base_url)
.finish()
}
}
impl ProviderChoice {
pub fn is_compatible_endpoint(&self) -> bool {
self.base_url.is_some()
}
}
#[derive(Debug, Error)]
pub enum ProviderError {
#[error(
"no provider credential found; set one of: {}",
CANDIDATES.iter().map(|(_, var)| *var).collect::<Vec<_>>().join(", ")
)]
NoCredential,
#[error("{provider} selected but {var} is not set")]
MissingCredential {
provider: BuiltinProvider,
var: &'static str,
},
#[error(
"unknown provider '{0}'; expected one of: anthropic, openai, gemini, openrouter, ollama, lmstudio"
)]
Unknown(String),
#[error("{0} has no API-key environment variable; it is a local provider")]
NotKeyed(BuiltinProvider),
#[error(
"a base URL was given but no key; set one of: {}",
COMPATIBLE_KEY_VARS.join(", ")
)]
NoCompatibleCredential,
#[error("base URL must be an absolute http(s) URL, got '{0}'")]
InvalidBaseUrl(String),
#[error("an API key was supplied with no provider and no base URL to attribute it to")]
UnattributedCredential,
}
pub fn normalize_base_url(raw: &str) -> Result<String, ProviderError> {
let trimmed = raw.trim();
let rest = trimmed
.strip_prefix("http://")
.or_else(|| trimmed.strip_prefix("https://"))
.ok_or_else(|| ProviderError::InvalidBaseUrl(raw.to_string()))?;
let host = rest.split('/').next().unwrap_or_default();
if host.is_empty() {
return Err(ProviderError::InvalidBaseUrl(raw.to_string()));
}
let without_slash = trimmed.trim_end_matches('/');
let without_version = without_slash
.strip_suffix("/v1")
.unwrap_or(without_slash)
.trim_end_matches('/');
if without_version.is_empty() {
return Err(ProviderError::InvalidBaseUrl(raw.to_string()));
}
Ok(format!("{without_version}/"))
}
pub fn parse(name: &str) -> Result<BuiltinProvider, ProviderError> {
match name.trim().to_ascii_lowercase().as_str() {
"anthropic" => Ok(BuiltinProvider::Anthropic),
"openai" => Ok(BuiltinProvider::OpenAI),
"gemini" => Ok(BuiltinProvider::Gemini),
"openrouter" => Ok(BuiltinProvider::OpenRouter),
"ollama" => Ok(BuiltinProvider::Ollama),
"lmstudio" | "lm-studio" => Ok(BuiltinProvider::LmStudio),
other => Err(ProviderError::Unknown(other.to_string())),
}
}
pub fn key_var(provider: BuiltinProvider) -> Option<&'static str> {
CANDIDATES
.iter()
.find(|(candidate, _)| *candidate == provider)
.map(|(_, var)| *var)
}
pub fn resolve(
requested: Option<BuiltinProvider>,
base_url: Option<&str>,
) -> Result<ProviderChoice, ProviderError> {
resolve_with(requested, base_url, None)
}
pub fn resolve_with(
requested: Option<BuiltinProvider>,
base_url: Option<&str>,
api_key: Option<&str>,
) -> Result<ProviderChoice, ProviderError> {
resolve_against(&|var| std::env::var(var).ok(), requested, base_url, api_key)
}
fn resolve_against(
lookup: &dyn Fn(&str) -> Option<String>,
requested: Option<BuiltinProvider>,
base_url: Option<&str>,
api_key: Option<&str>,
) -> Result<ProviderChoice, ProviderError> {
if let Some(raw) = base_url
.map(str::to_string)
.or_else(|| env_base_url(lookup))
{
return resolve_compatible(lookup, &raw, requested, api_key);
}
match (requested, api_key) {
(Some(provider), Some(api_key)) => Ok(ProviderChoice {
provider,
api_key: api_key.to_string(),
source_var: None,
base_url: None,
}),
(None, Some(_)) => Err(ProviderError::UnattributedCredential),
(Some(provider), None) => {
let var = key_var(provider).ok_or(ProviderError::NotKeyed(provider))?;
let api_key =
read(lookup, var).ok_or(ProviderError::MissingCredential { provider, var })?;
Ok(ProviderChoice {
provider,
api_key,
source_var: Some(var),
base_url: None,
})
}
(None, None) => CANDIDATES
.iter()
.find_map(|(provider, var)| {
read(lookup, var).map(|api_key| ProviderChoice {
provider: *provider,
api_key,
source_var: Some(var),
base_url: None,
})
})
.ok_or(ProviderError::NoCredential),
}
}
fn resolve_compatible(
lookup: &dyn Fn(&str) -> Option<String>,
raw: &str,
requested: Option<BuiltinProvider>,
api_key: Option<&str>,
) -> Result<ProviderChoice, ProviderError> {
let base_url = normalize_base_url(raw)?;
let (api_key, source_var) = match api_key {
Some(api_key) => (api_key.to_string(), None),
None => COMPATIBLE_KEY_VARS
.iter()
.find_map(|var| read(lookup, var).map(|key| (key, Some(*var))))
.ok_or(ProviderError::NoCompatibleCredential)?,
};
Ok(ProviderChoice {
provider: requested.unwrap_or(BuiltinProvider::OpenAI),
api_key,
source_var,
base_url: Some(base_url),
})
}
fn env_base_url(lookup: &dyn Fn(&str) -> Option<String>) -> Option<String> {
BASE_URL_VARS.iter().find_map(|var| read(lookup, var))
}
fn read(lookup: &dyn Fn(&str) -> Option<String>, var: &str) -> Option<String> {
lookup(var).filter(|value| !value.trim().is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
fn exporting(vars: &[(&str, &str)]) -> impl Fn(&str) -> Option<String> {
let vars: Vec<(String, String)> = vars
.iter()
.map(|(var, value)| (var.to_string(), value.to_string()))
.collect();
move |name| {
vars.iter()
.find(|(var, _)| var == name)
.map(|(_, value)| value.clone())
}
}
fn nothing_exported() -> impl Fn(&str) -> Option<String> {
exporting(&[])
}
#[test]
fn provider_names_parse_case_insensitively() {
assert_eq!(parse("OpenAI").expect("parses"), BuiltinProvider::OpenAI);
assert_eq!(
parse(" anthropic ").expect("parses"),
BuiltinProvider::Anthropic
);
assert_eq!(
parse("lm-studio").expect("parses"),
BuiltinProvider::LmStudio
);
}
#[test]
fn an_unknown_provider_names_the_alternatives() {
let error = parse("hal9000").expect_err("rejected");
assert!(matches!(error, ProviderError::Unknown(name) if name == "hal9000"));
}
#[test]
fn hosted_providers_have_a_key_variable_and_local_ones_do_not() {
assert_eq!(
key_var(BuiltinProvider::Anthropic),
Some("ANTHROPIC_API_KEY")
);
assert_eq!(key_var(BuiltinProvider::Ollama), None);
}
#[test]
fn detection_order_prefers_the_first_candidate() {
let vars: Vec<&str> = CANDIDATES.iter().map(|(_, var)| *var).collect();
assert_eq!(vars.first(), Some(&"ANTHROPIC_API_KEY"));
assert_eq!(
vars.len(),
4,
"local providers must not be auto-detection candidates"
);
}
#[test]
fn detection_takes_the_first_candidate_the_environment_offers() {
let choice = resolve_against(
&exporting(&[
("OPENAI_API_KEY", "openai-key"),
("ANTHROPIC_API_KEY", "anthropic-key"),
]),
None,
None,
None,
)
.expect("a key is exported");
assert_eq!(choice.provider, BuiltinProvider::Anthropic);
assert_eq!(choice.source_var, Some("ANTHROPIC_API_KEY"));
}
#[test]
fn a_named_provider_reads_its_own_variable_and_says_which() {
let choice = resolve_against(
&exporting(&[
("ANTHROPIC_API_KEY", "anthropic-key"),
("GEMINI_API_KEY", "gemini-key"),
]),
Some(BuiltinProvider::Gemini),
None,
None,
)
.expect("the named provider's key is exported");
assert_eq!(choice.api_key, "gemini-key");
assert_eq!(choice.source_var, Some("GEMINI_API_KEY"));
}
#[test]
fn a_variable_set_to_whitespace_is_treated_as_absent() {
let error = resolve_against(
&exporting(&[("ANTHROPIC_API_KEY", " ")]),
None,
None,
None,
)
.expect_err("rejected");
assert!(matches!(error, ProviderError::NoCredential));
}
#[test]
fn an_environment_base_url_outranks_provider_detection() {
let choice = resolve_against(
&exporting(&[
("ANTHROPIC_API_KEY", "anthropic-key"),
("BASIS_BASE_URL", "http://127.0.0.1:3455/v1"),
("BASIS_API_KEY", "gateway-key"),
]),
None,
None,
None,
)
.expect("a base URL and a key are enough");
assert_eq!(choice.base_url.as_deref(), Some("http://127.0.0.1:3455/"));
assert_eq!(choice.api_key, "gateway-key");
assert_eq!(choice.source_var, Some("BASIS_API_KEY"));
}
#[test]
fn a_base_url_with_no_key_anywhere_is_refused() {
let error = resolve_against(
&exporting(&[("BASIS_BASE_URL", "http://127.0.0.1:3455/v1")]),
None,
None,
None,
)
.expect_err("rejected");
assert!(matches!(error, ProviderError::NoCompatibleCredential));
}
#[test]
fn selecting_a_local_provider_by_key_is_rejected() {
let error = resolve_against(
¬hing_exported(),
Some(BuiltinProvider::Ollama),
None,
None,
)
.expect_err("rejected");
assert!(matches!(error, ProviderError::NotKeyed(_)));
}
#[test]
fn a_named_provider_with_no_key_names_the_variable_it_wanted() {
let error = resolve_against(
¬hing_exported(),
Some(BuiltinProvider::OpenRouter),
None,
None,
)
.expect_err("rejected");
assert!(matches!(
error,
ProviderError::MissingCredential {
var: "OPENROUTER_API_KEY",
..
}
));
}
#[test]
fn a_supplied_key_is_used_instead_of_the_environment() {
let choice = resolve_against(
&exporting(&[("ANTHROPIC_API_KEY", "exported-key")]),
Some(BuiltinProvider::Anthropic),
None,
Some("supplied-key"),
)
.expect("a named provider and a key need no lookup");
assert_eq!(choice.api_key, "supplied-key");
assert_eq!(choice.provider, BuiltinProvider::Anthropic);
assert_eq!(
choice.source_var, None,
"no variable was read, so none may be named"
);
}
#[test]
fn a_supplied_key_reaches_a_compatible_endpoint() {
let choice = resolve_against(
¬hing_exported(),
None,
Some("http://127.0.0.1:3455/v1"),
Some("supplied-key"),
)
.expect("a base URL and a key are enough");
assert_eq!(choice.api_key, "supplied-key");
assert_eq!(choice.base_url.as_deref(), Some("http://127.0.0.1:3455/"));
assert!(choice.is_compatible_endpoint());
}
#[test]
fn a_key_with_nothing_to_attribute_it_to_is_refused() {
let error = resolve_against(¬hing_exported(), None, None, Some("supplied-key"))
.expect_err("rejected");
assert!(matches!(error, ProviderError::UnattributedCredential));
}
#[test]
fn a_resolved_credential_is_not_printed() {
let choice = resolve_against(
&exporting(&[("ANTHROPIC_API_KEY", "sk-secret-value")]),
None,
None,
None,
)
.expect("a key is exported");
let printed = format!("{choice:?}");
assert!(!printed.contains("sk-secret-value"));
assert!(printed.contains("redacted"));
assert!(
printed.contains("ANTHROPIC_API_KEY"),
"which variable answered is not the secret, and is how a caller debugs this"
);
}
#[test]
fn a_published_base_url_keeps_its_host_and_loses_its_version_suffix() {
assert_eq!(
normalize_base_url("http://127.0.0.1:3455/v1").expect("normalizes"),
"http://127.0.0.1:3455/"
);
assert_eq!(
normalize_base_url("https://gateway.example.com/v1/").expect("normalizes"),
"https://gateway.example.com/"
);
}
#[test]
fn a_base_url_without_a_version_suffix_is_left_alone() {
assert_eq!(
normalize_base_url("https://gateway.example.com").expect("normalizes"),
"https://gateway.example.com/"
);
}
#[test]
fn a_path_prefix_survives_normalization() {
assert_eq!(
normalize_base_url("https://example.com/openai/v1").expect("normalizes"),
"https://example.com/openai/"
);
}
#[test]
fn a_base_url_must_be_absolute_http() {
for raw in ["127.0.0.1:3455/v1", "ftp://example.com", "", "https://"] {
assert!(
normalize_base_url(raw).is_err(),
"'{raw}' must be rejected before it reaches the transport"
);
}
}
#[test]
fn an_endpoint_is_flagged_as_compatible() {
let choice = ProviderChoice {
provider: BuiltinProvider::OpenAI,
api_key: "k".to_string(),
source_var: None,
base_url: Some("http://localhost:1/".to_string()),
};
assert!(choice.is_compatible_endpoint());
}
}