use async_trait::async_trait;
use serde_json::Value;
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::{ProviderCapabilities, ProviderId};
use crate::inference::streaming::ChatStream;
use crate::inference::types::{ChatRequest, ChatResponse, SecretString, ToolChoice};
use crate::local_probe::{self, LocalProbeError};
pub const LOCAL_BASE_URL: &str = "http://localhost:11434/v1";
pub use crate::local_probe::LOCAL_HOST_ENV;
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 = format!("{}/v1", local_probe::local_host());
let auth = std::env::var(LOCAL_API_KEY_ENV)
.ok()
.filter(|v| !v.trim().is_empty())
.map(SecretString::new);
Self { base_url, auth }
}
}
fn probe_failure(err: LocalProbeError) -> InferenceError {
InferenceError::Transport(err.to_string())
}
pub struct LocalAdapter {
inner: OpenAiCompatAdapter,
probe_url: String,
}
impl LocalAdapter {
pub fn new(inner: OpenAiCompatAdapter, base_url: &str) -> Self {
Self {
inner,
probe_url: local_probe::models_url(base_url),
}
}
pub fn probe_url(&self) -> &str {
&self.probe_url
}
pub async fn probe(&self) -> Result<(), InferenceError> {
local_probe::probe_models_endpoint(&self.probe_url)
.await
.map_err(probe_failure)
}
}
#[async_trait]
impl InferenceAdapter for LocalAdapter {
fn name(&self) -> &str {
self.inner.name()
}
fn capabilities(&self) -> &ProviderCapabilities {
self.inner.capabilities()
}
fn capabilities_for(&self, model: &str) -> &ProviderCapabilities {
self.inner.capabilities_for(model)
}
async fn chat(&self, request: &ChatRequest) -> Result<ChatResponse, InferenceError> {
self.ensure_structured_output_supported(request)?;
self.probe().await?;
self.inner.chat(request).await
}
async fn chat_stream(&self, request: &ChatRequest) -> Result<ChatStream, InferenceError> {
self.ensure_structured_output_supported(request)?;
self.probe().await?;
self.inner.chat_stream(request).await
}
fn map_tool_choice(&self, choice: ToolChoice) -> Value {
self.inner.map_tool_choice(choice)
}
fn supports_native_tools(&self) -> bool {
self.inner.supports_native_tools()
}
fn supports_prompt_caching(&self) -> bool {
self.inner.supports_prompt_caching()
}
fn supports_structured_output(&self) -> bool {
self.inner.supports_structured_output()
}
fn wants_detailed_usage(&self) -> bool {
self.inner.wants_detailed_usage()
}
fn context_window(&self, model: &str) -> usize {
self.inner.context_window(model)
}
}
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 base_url = config.base_url;
let cfg = OpenAiCompatConfig {
name: ProviderId::Local.as_str().to_string(),
base_url: base_url.clone(),
api_key,
extra_headers: Vec::new(),
capabilities: *crate::inference::registry::capabilities(ProviderId::Local),
};
let inner = OpenAiCompatAdapter::new(cfg)?;
Ok(Box::new(LocalAdapter::new(inner, &base_url)))
}
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");
}
#[test]
fn probe_url_is_derived_from_the_base_url() {
let inner = OpenAiCompatAdapter::new(OpenAiCompatConfig {
name: ProviderId::Local.as_str().to_string(),
base_url: LOCAL_BASE_URL.to_string(),
api_key: SecretString::new(LOCAL_PLACEHOLDER_KEY),
extra_headers: Vec::new(),
capabilities: *crate::inference::registry::capabilities(ProviderId::Local),
})
.expect("inner adapter builds");
let adapter = LocalAdapter::new(inner, LOCAL_BASE_URL);
assert_eq!(adapter.probe_url(), "http://localhost:11434/v1/models");
}
#[tokio::test]
async fn chat_fails_inside_the_probe_budget_when_the_server_never_answers() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind black-hole listener");
let addr = listener.local_addr().expect("listener addr");
tokio::spawn(async move {
let mut held = Vec::new();
while let Ok((stream, _)) = listener.accept().await {
held.push(stream);
}
});
let config = LocalConfig {
base_url: format!("http://{addr}/v1"),
auth: None,
};
let adapter = build(&resolved(), config).expect("built");
let request = ChatRequest::new(
"local/llama3.1",
vec![crate::inference::types::ChatMessage::user("ping")],
);
let started = std::time::Instant::now();
let outcome =
tokio::time::timeout(std::time::Duration::from_secs(5), adapter.chat(&request))
.await
.expect(
"chat must return inside the probe budget, not hang on a dead local server",
);
let elapsed = started.elapsed();
let err = outcome.expect_err("a server that never answers must not yield a response");
assert!(
matches!(err, InferenceError::Transport(_)),
"expected Transport, got {err:?}"
);
assert!(
err.to_string().contains(&addr.to_string()),
"the error must name the endpoint that was dialled: {err}"
);
assert!(err.is_retryable(), "a dead local server is worth retrying");
assert!(
elapsed < std::time::Duration::from_secs(3),
"probe budget is {:?}; took {elapsed:?}",
crate::local_probe::LOCAL_PROBE_TIMEOUT
);
}
}