use std::time::Duration;
pub const LOCAL_PROBE_TIMEOUT: Duration = Duration::from_secs(1);
pub const LOCAL_MODELS_PATH: &str = "/v1/models";
pub const LOCAL_HOST_ENV: &str = "OLLAMA_HOST";
pub const DEFAULT_LOCAL_HOST: &str = "http://localhost:11434";
pub fn local_host() -> String {
std::env::var(LOCAL_HOST_ENV)
.ok()
.filter(|v| !v.trim().is_empty())
.map(|host| host.trim_end_matches('/').to_string())
.unwrap_or_else(|| DEFAULT_LOCAL_HOST.to_string())
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum LocalProbeError {
ClientBuild {
endpoint: String,
cause: String,
},
Unreachable {
endpoint: String,
cause: String,
},
Status {
endpoint: String,
status: u16,
},
Body {
endpoint: String,
cause: String,
},
}
impl LocalProbeError {
pub fn endpoint(&self) -> &str {
match self {
Self::ClientBuild { endpoint, .. }
| Self::Unreachable { endpoint, .. }
| Self::Status { endpoint, .. }
| Self::Body { endpoint, .. } => endpoint,
}
}
}
impl std::fmt::Display for LocalProbeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::ClientBuild { endpoint, cause } => write!(
f,
"could not build a probe client for local model server {endpoint}: {cause}"
),
Self::Unreachable { endpoint, cause } => write!(
f,
"local model server not reachable at {endpoint} within {}s: {cause}",
LOCAL_PROBE_TIMEOUT.as_secs()
),
Self::Status { endpoint, status } => write!(
f,
"local model server at {endpoint} answered HTTP {status}, not a success status"
),
Self::Body { endpoint, cause } => write!(
f,
"local model server at {endpoint} answered with an unreadable model list: {cause}"
),
}
}
}
impl std::error::Error for LocalProbeError {}
pub fn models_url(base_url: &str) -> String {
let base = base_url.trim_end_matches('/');
let host = base.strip_suffix("/v1").unwrap_or(base);
format!("{host}{LOCAL_MODELS_PATH}")
}
pub async fn probe_models_endpoint(url: &str) -> Result<(), LocalProbeError> {
get_success(url).await.map(|_| ())
}
async fn get_success(url: &str) -> Result<reqwest::Response, LocalProbeError> {
let client = reqwest::Client::builder()
.connect_timeout(LOCAL_PROBE_TIMEOUT)
.timeout(LOCAL_PROBE_TIMEOUT)
.build()
.map_err(|e| LocalProbeError::ClientBuild {
endpoint: url.to_string(),
cause: e.to_string(),
})?;
match client.get(url).send().await {
Ok(resp) if resp.status().is_success() => Ok(resp),
Ok(resp) => Err(LocalProbeError::Status {
endpoint: url.to_string(),
status: resp.status().as_u16(),
}),
Err(e) => Err(LocalProbeError::Unreachable {
endpoint: url.to_string(),
cause: e.to_string(),
}),
}
}
pub async fn probe_local(base_url: &str) -> Result<(), LocalProbeError> {
probe_models_endpoint(&models_url(base_url)).await
}
pub async fn list_models(base_url: &str) -> Result<Vec<String>, LocalProbeError> {
let url = models_url(base_url);
let resp = get_success(&url).await?;
let body: serde_json::Value = resp.json().await.map_err(|e| LocalProbeError::Body {
endpoint: url.clone(),
cause: e.to_string(),
})?;
Ok(body
.get("data")
.and_then(|v| v.as_array())
.map(|arr| {
arr.iter()
.filter_map(|m| m.get("id").and_then(|id| id.as_str()).map(str::to_string))
.collect()
})
.unwrap_or_default())
}
#[cfg(test)]
mod tests {
use super::*;
async fn stub_server(response: impl Into<String>) -> String {
let response: std::sync::Arc<str> = response.into().into();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind loopback stub");
let addr = listener.local_addr().expect("stub addr").to_string();
tokio::spawn(async move {
while let Ok((mut stream, _)) = listener.accept().await {
let response = std::sync::Arc::clone(&response);
tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
let _ = stream.write_all(response.as_bytes()).await;
});
}
});
addr
}
fn json_ok(body: &str) -> String {
format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{body}",
body.len()
)
}
fn dead_addr() -> String {
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind to free a port");
let addr = listener.local_addr().expect("dead addr").to_string();
drop(listener);
addr
}
#[test]
fn models_url_appends_v1_when_absent() {
assert_eq!(
models_url("http://localhost:11434"),
"http://localhost:11434/v1/models"
);
assert_eq!(
models_url("http://localhost:11434/"),
"http://localhost:11434/v1/models"
);
}
#[test]
fn models_url_does_not_double_an_existing_v1_suffix() {
assert_eq!(
models_url("http://localhost:11434/v1"),
"http://localhost:11434/v1/models"
);
assert_eq!(
models_url("http://localhost:11434/v1/"),
"http://localhost:11434/v1/models"
);
}
#[tokio::test]
async fn probe_reports_unreachable_naming_the_endpoint() {
let addr = dead_addr();
let base = format!("http://{addr}");
let err = probe_local(&base).await.expect_err("closed port must fail");
assert!(
matches!(err, LocalProbeError::Unreachable { .. }),
"expected Unreachable, got {err:?}"
);
assert_eq!(err.endpoint(), format!("{base}/v1/models"));
assert!(err.to_string().contains(&addr), "{err}");
}
#[tokio::test]
async fn probe_reports_non_success_status() {
let addr = stub_server("HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n").await;
let err = probe_local(&format!("http://{addr}"))
.await
.expect_err("404 must fail");
assert_eq!(
err,
LocalProbeError::Status {
endpoint: format!("http://{addr}/v1/models"),
status: 404,
}
);
}
#[tokio::test]
async fn probe_accepts_a_live_endpoint() {
let addr = stub_server("HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok").await;
probe_local(&format!("http://{addr}"))
.await
.expect("live endpoint must probe clean");
}
#[test]
fn probe_timeout_is_one_second() {
assert_eq!(LOCAL_PROBE_TIMEOUT, Duration::from_secs(1));
}
#[tokio::test]
async fn list_models_returns_the_served_ids() {
let addr = stub_server(json_ok(
r#"{"object":"list","data":[{"id":"qwen3:30b"},{"id":"llama3.1:8b"}]}"#,
))
.await;
let models = list_models(&format!("http://{addr}"))
.await
.expect("live endpoint must list");
assert_eq!(models, vec!["qwen3:30b", "llama3.1:8b"]);
}
#[tokio::test]
async fn list_models_reports_an_unreadable_body() {
let addr = stub_server(
"HTTP/1.1 200 OK\r\nContent-Type: text/html\r\nContent-Length: 5\r\n\r\nnope!",
)
.await;
let base = format!("http://{addr}");
let err = list_models(&base).await.expect_err("non-JSON must fail");
assert!(
matches!(err, LocalProbeError::Body { .. }),
"expected Body, got {err:?}"
);
assert_eq!(err.endpoint(), format!("{base}/v1/models"));
}
#[tokio::test]
async fn list_models_is_empty_when_the_server_serves_none() {
let addr = stub_server(json_ok(r#"{"object":"list","data":[]}"#)).await;
assert!(
list_models(&format!("http://{addr}"))
.await
.expect("empty catalog is still live")
.is_empty()
);
}
#[test]
#[serial_test::serial]
fn local_host_reads_the_env_override() {
unsafe { std::env::set_var(LOCAL_HOST_ENV, "http://192.168.1.50:11434/") };
assert_eq!(local_host(), "http://192.168.1.50:11434");
unsafe { std::env::remove_var(LOCAL_HOST_ENV) };
}
#[test]
#[serial_test::serial]
fn local_host_defaults_when_unset() {
unsafe { std::env::remove_var(LOCAL_HOST_ENV) };
assert_eq!(local_host(), DEFAULT_LOCAL_HOST);
unsafe { std::env::set_var(LOCAL_HOST_ENV, " ") };
assert_eq!(local_host(), DEFAULT_LOCAL_HOST);
unsafe { std::env::remove_var(LOCAL_HOST_ENV) };
}
}