use super::openai_compat::{OpenAiCompatAdapter, OpenAiCompatConfig};
use crate::inference::adapter::InferenceAdapter;
use crate::inference::configurator::ResolvedProvider;
use crate::inference::error::InferenceError;
use crate::inference::registry::ProviderId;
use crate::inference::types::SecretString;
pub const LOCAL_BASE_URL: &str = "http://localhost:11434/v1";
pub const LOCAL_HOST_ENV: &str = "OLLAMA_HOST";
pub const LOCAL_API_KEY_ENV: &str = "TRUSTY_LOCAL_API_KEY";
pub const LOCAL_PLACEHOLDER_KEY: &str = "not-needed";
pub struct LocalConfig {
pub base_url: String,
pub auth: Option<SecretString>,
}
impl LocalConfig {
pub fn from_env() -> Self {
let base_url = std::env::var(LOCAL_HOST_ENV)
.ok()
.filter(|v| !v.trim().is_empty())
.map(|host| format!("{}/v1", host.trim_end_matches('/')))
.unwrap_or_else(|| LOCAL_BASE_URL.to_string());
let auth = std::env::var(LOCAL_API_KEY_ENV)
.ok()
.filter(|v| !v.trim().is_empty())
.map(SecretString::new);
Self { base_url, auth }
}
}
pub fn build(
_resolved: &ResolvedProvider,
config: LocalConfig,
) -> Result<Box<dyn InferenceAdapter>, InferenceError> {
let api_key = config
.auth
.unwrap_or_else(|| SecretString::new(LOCAL_PLACEHOLDER_KEY));
let cfg = OpenAiCompatConfig {
name: ProviderId::Local.as_str().to_string(),
base_url: config.base_url,
api_key,
extra_headers: Vec::new(),
capabilities: *crate::inference::registry::capabilities(ProviderId::Local),
};
Ok(Box::new(OpenAiCompatAdapter::new(cfg)?))
}
pub fn factory(resolved: &ResolvedProvider) -> Result<Box<dyn InferenceAdapter>, InferenceError> {
build(resolved, LocalConfig::from_env())
}
#[cfg(test)]
mod tests {
use super::*;
use serial_test::serial;
fn resolved() -> ResolvedProvider {
ResolvedProvider::new(ProviderId::Local, "local/llama3.1".to_string(), None)
}
#[test]
#[serial(local_provider_env)]
fn factory_builds_named_adapter_with_defaults() {
unsafe {
std::env::remove_var(LOCAL_HOST_ENV);
std::env::remove_var(LOCAL_API_KEY_ENV);
}
let adapter = build(&resolved(), LocalConfig::from_env()).expect("built");
assert_eq!(adapter.name(), "local");
assert!(adapter.supports_native_tools());
assert_eq!(adapter.capabilities().id, ProviderId::Local);
}
#[test]
#[serial(local_provider_env)]
fn from_env_defaults_when_unset() {
unsafe {
std::env::remove_var(LOCAL_HOST_ENV);
std::env::remove_var(LOCAL_API_KEY_ENV);
}
let config = LocalConfig::from_env();
assert_eq!(config.base_url, LOCAL_BASE_URL);
assert!(config.auth.is_none());
}
#[test]
#[serial(local_provider_env)]
fn host_env_override_appends_v1_suffix() {
unsafe {
std::env::set_var(LOCAL_HOST_ENV, "http://192.168.1.50:11434/");
std::env::remove_var(LOCAL_API_KEY_ENV);
}
let config = LocalConfig::from_env();
assert_eq!(config.base_url, "http://192.168.1.50:11434/v1");
let adapter = build(&resolved(), config).expect("built");
assert_eq!(adapter.capabilities().id.credential_name(), None);
unsafe {
std::env::remove_var(LOCAL_HOST_ENV);
}
}
#[test]
#[serial(local_provider_env)]
fn api_key_env_override_is_used() {
unsafe {
std::env::remove_var(LOCAL_HOST_ENV);
std::env::set_var(LOCAL_API_KEY_ENV, "sk-local-test"); }
let config = LocalConfig::from_env();
assert_eq!(
config.auth.as_ref().map(SecretString::expose),
Some("sk-local-test")
);
unsafe {
std::env::remove_var(LOCAL_API_KEY_ENV);
}
}
#[test]
fn placeholder_key_used_when_no_override() {
let config = LocalConfig {
base_url: LOCAL_BASE_URL.to_string(),
auth: None,
};
let adapter = build(&resolved(), config).expect("built without any credential");
assert_eq!(adapter.name(), "local");
}
}