use crate::models::huggingface::hf_repo_id;
use crate::models::{ModelFunction, ModelMetadata, ModelVariant};
use crate::providers::base::{
ApiEndpoint, ApiType, AuthType, HasProviderMetadata, HealthStatus, ModelFormat, Provider,
ProviderError, ProviderMetadata, ProviderType,
};
use crate::registry::{ConfigConstructable, Secret};
use crate::utils::ui::Ui;
use crate::utils::ui::base::PullHandle;
use async_trait::async_trait;
use futures_util::StreamExt;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::time::Duration;
#[derive(Debug, Deserialize)]
struct LlamaCppSseFrame {
model: String,
event: String,
#[serde(default)]
data: serde_json::Value,
}
#[derive(Debug, Deserialize)]
struct LlamaCppModelsResponse {
data: Vec<LlamaCppModelEntry>,
}
#[derive(Debug, Deserialize)]
struct LlamaCppModelEntry {
id: String,
status: LlamaCppModelStatus,
}
#[derive(Debug, Deserialize)]
struct LlamaCppModelStatus {
value: String,
#[serde(default)]
failed: bool,
}
fn find_double_newline(buf: &[u8]) -> Option<usize> {
buf.windows(2).position(|w| w == b"\n\n")
}
#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
pub struct LlamaCppProviderConfig {
#[serde(default = "default_llamacpp_url")]
pub base_url: String,
pub api_key: Option<Secret>,
#[serde(default = "default_timeout")]
pub timeout_secs: u64,
#[serde(default = "default_verify_ssl")]
pub verify_ssl: bool,
#[serde(default = "default_llamacpp_health_endpoint")]
pub health_check_endpoint: String,
}
fn default_llamacpp_url() -> String {
"http://localhost:8080".to_string()
}
fn default_timeout() -> u64 {
10
}
fn default_verify_ssl() -> bool {
true
}
fn default_llamacpp_health_endpoint() -> String {
"/health".to_string()
}
impl Default for LlamaCppProviderConfig {
fn default() -> Self {
Self {
base_url: default_llamacpp_url(),
api_key: None,
timeout_secs: default_timeout(),
verify_ssl: default_verify_ssl(),
health_check_endpoint: default_llamacpp_health_endpoint(),
}
}
}
pub struct LlamaCppProvider {
instance_id: String,
config: LlamaCppProviderConfig,
client: reqwest::Client,
stream_client: reqwest::Client,
}
impl LlamaCppProvider {
fn default_function_endpoints() -> HashMap<ModelFunction, Vec<ApiEndpoint>> {
let mut map = HashMap::new();
map.insert(
ModelFunction::Chat,
vec![ApiEndpoint::OpenAIChat, ApiEndpoint::AnthropicMessages],
);
map.insert(
ModelFunction::ToolCalling,
vec![ApiEndpoint::OpenAIChat, ApiEndpoint::AnthropicMessages],
);
map.insert(
ModelFunction::Thinking,
vec![ApiEndpoint::OpenAIChat, ApiEndpoint::AnthropicMessages],
);
map.insert(
ModelFunction::ImageUnderstanding,
vec![ApiEndpoint::OpenAIChat, ApiEndpoint::AnthropicMessages],
);
map.insert(
ModelFunction::Guardian,
vec![ApiEndpoint::OpenAIChat, ApiEndpoint::AnthropicMessages],
);
map.insert(
ModelFunction::Embeddings,
vec![ApiEndpoint::OpenAIEmbeddings],
);
map.insert(
ModelFunction::Transcription,
vec![ApiEndpoint::OpenAIAudioTranscription],
);
map
}
async fn watch_pull_via_sse(
&self,
model_ref: &str,
handle: PullHandle,
label: &str,
ui: &dyn Ui,
) -> Result<bool, ProviderError> {
let sse_url = format!("{}/models/sse", self.config.base_url);
let mut request = self.stream_client.get(&sse_url);
if let Some(key) = &self.config.api_key {
request = request.bearer_auth(&key.0);
}
let response = match request.send().await {
Ok(r) if r.status().is_success() => r,
_ => return Ok(false),
};
let mut stream = response.bytes_stream();
let mut buf: Vec<u8> = Vec::new();
while let Some(chunk) = stream.next().await {
let chunk = match chunk {
Ok(c) => c,
Err(_) => return Ok(false),
};
buf.extend_from_slice(&chunk);
while let Some(pos) = find_double_newline(&buf) {
let frame_bytes: Vec<u8> = buf.drain(..pos + 2).collect();
let frame_text = String::from_utf8_lossy(&frame_bytes);
let Some(json_str) = frame_text.trim().strip_prefix("data:") else {
continue;
};
let json_str = json_str.trim();
if json_str.is_empty() {
continue;
}
let frame: LlamaCppSseFrame = match serde_json::from_str(json_str) {
Ok(f) => f,
Err(_) => continue,
};
if frame.model != model_ref {
continue;
}
match frame.event.as_str() {
"download_progress" => {
if let Some(progress) =
frame.data.get("progress").and_then(|p| p.as_object())
{
let mut done_sum = 0u64;
let mut total_sum = 0u64;
for entry in progress.values() {
done_sum += entry.get("done").and_then(|v| v.as_u64()).unwrap_or(0);
total_sum +=
entry.get("total").and_then(|v| v.as_u64()).unwrap_or(0);
}
ui.pull_progress(
handle,
done_sum,
if total_sum > 0 { Some(total_sum) } else { None },
);
}
}
"download_finished" => {
ui.pull_finish(handle, label, None);
return Ok(true);
}
"download_failed" => {
let err = "download failed".to_string();
ui.pull_finish(handle, label, Some(&err));
return Err(ProviderError::Other(err));
}
_ => {}
}
}
}
Ok(false)
}
async fn watch_pull_via_polling(
&self,
repo: &str,
handle: PullHandle,
label: &str,
ui: &dyn Ui,
) -> Result<crate::providers::PullResult, ProviderError> {
let models_url = format!("{}/models", self.config.base_url);
loop {
tokio::time::sleep(Duration::from_secs(1)).await;
let mut request = self.client.get(&models_url);
if let Some(key) = &self.config.api_key {
request = request.bearer_auth(&key.0);
}
let response = request.send().await?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
let err = format!("llama.cpp models status check failed ({status}): {body}");
ui.pull_finish(handle, label, Some(&err));
return Err(ProviderError::Other(err));
}
let parsed: LlamaCppModelsResponse = response.json().await?;
let Some(entry) = parsed.data.iter().find(|e| e.id.contains(repo)) else {
continue;
};
if entry.status.failed {
let err = "download failed".to_string();
ui.pull_finish(handle, label, Some(&err));
return Err(ProviderError::Other(err));
}
if entry.status.value != "downloading" {
ui.pull_finish(handle, label, None);
return Ok(crate::providers::PullResult::Success);
}
}
}
}
impl ConfigConstructable for LlamaCppProvider {
type Config = LlamaCppProviderConfig;
fn new(
instance_id: &str,
cfg: &serde_json::Value,
_global_config: &crate::config::Config,
) -> Self {
let config: LlamaCppProviderConfig =
serde_json::from_value(cfg.clone()).unwrap_or_default();
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(config.timeout_secs))
.danger_accept_invalid_certs(!config.verify_ssl)
.build()
.expect("Failed to create HTTP client");
let stream_client = reqwest::Client::builder()
.danger_accept_invalid_certs(!config.verify_ssl)
.build()
.expect("Failed to create HTTP client");
Self {
instance_id: instance_id.to_string(),
config,
client,
stream_client,
}
}
}
impl crate::registry::Named for LlamaCppProvider {
fn instance_id(&self) -> &str {
&self.instance_id
}
}
#[async_trait]
impl Provider for LlamaCppProvider {
fn name(&self) -> &str {
"llama.cpp"
}
fn function_endpoints(&self) -> HashMap<ModelFunction, Vec<ApiEndpoint>> {
Self::default_function_endpoints()
}
fn supported_api_types(&self) -> Vec<ApiType> {
vec![ApiType::OpenAI, ApiType::Anthropic]
}
fn base_url(&self) -> &str {
&self.config.base_url
}
fn api_key(&self) -> Option<&Secret> {
self.config.api_key.as_ref()
}
fn verify_ssl(&self) -> bool {
self.config.verify_ssl
}
fn supported_formats(&self) -> Vec<ModelFormat> {
vec![ModelFormat::GGUF]
}
fn can_run_model(&self, variant_format: &str, _variant_precision: &str) -> bool {
variant_format.eq_ignore_ascii_case("gguf")
}
fn model_alias(
&self,
_model_id: String,
variant: Option<&crate::models::ModelVariant>,
) -> Option<String> {
let v = variant?;
let repo = hf_repo_id(&v.url)?;
Some(format!("{}:{}", repo, v.precision))
}
async fn health_check(&self) -> Result<HealthStatus, ProviderError> {
use std::time::Instant;
let start = Instant::now();
let url = format!(
"{}{}",
self.config.base_url, self.config.health_check_endpoint
);
let mut request = self.client.get(&url);
if let Some(key) = &self.config.api_key {
request = request.bearer_auth(&key.0);
}
match request.send().await {
Ok(response) => {
let latency = start.elapsed();
if response.status() != reqwest::StatusCode::OK {
return Ok(HealthStatus {
healthy: false,
latency,
error: Some(format!(
"HTTP {}: {}",
response.status(),
response.text().await.unwrap_or_default()
)),
});
}
match response.json::<serde_json::Value>().await {
Ok(body) => {
let status = body.get("status").and_then(|s| s.as_str()).unwrap_or("");
if status == "ok" {
Ok(HealthStatus {
healthy: true,
latency,
error: None,
})
} else {
Ok(HealthStatus {
healthy: false,
latency,
error: Some(format!("unexpected status: {status:?}")),
})
}
}
Err(e) => Ok(HealthStatus {
healthy: false,
latency,
error: Some(format!("invalid JSON response: {e}")),
}),
}
}
Err(e) => {
let latency = start.elapsed();
Ok(HealthStatus {
healthy: false,
latency,
error: Some(format!("Connection failed: {e}")),
})
}
}
}
async fn pull_model(
&self,
model: &ModelMetadata,
variant: &ModelVariant,
ui: &dyn Ui,
) -> Result<crate::providers::PullResult, ProviderError> {
let repo = hf_repo_id(&variant.url).ok_or_else(|| {
ProviderError::Other(format!(
"cannot determine a HuggingFace repo for {} variant {}/{}",
model.family, variant.format, variant.precision
))
})?;
let model_ref = format!("{}:{}", repo, variant.precision);
let label = format!(
"{} ({} {})",
model.family, variant.format, variant.precision
);
let post_url = format!("{}/models", self.config.base_url);
let mut request = self
.client
.post(&post_url)
.json(&serde_json::json!({ "model": model_ref }));
if let Some(key) = &self.config.api_key {
request = request.bearer_auth(&key.0);
}
let response = request.send().await?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
if status == reqwest::StatusCode::BAD_REQUEST && body.contains("already exists") {
ui.info(&format!("{label} is already downloaded."));
return Ok(crate::providers::PullResult::Success);
}
return Err(ProviderError::Other(format!(
"llama.cpp model pull request failed ({status}): {body}"
)));
}
let handle = ui.pull_start(&label, None);
match self
.watch_pull_via_sse(&model_ref, handle, &label, ui)
.await
{
Ok(true) => Ok(crate::providers::PullResult::Success),
Ok(false) => self.watch_pull_via_polling(repo, handle, &label, ui).await,
Err(e) => Err(e),
}
}
}
impl HasProviderMetadata for LlamaCppProvider {
fn metadata() -> ProviderMetadata {
ProviderMetadata {
name: "llama.cpp".to_string(),
description: "High-performance local inference server for GGUF models with OpenAI and Anthropic API compatibility".to_string(),
provider_type: ProviderType::Local,
default_endpoint: "http://localhost:8080".to_string(),
supported_api_types: vec![ApiType::OpenAI, ApiType::Anthropic],
default_function_endpoints: Self::default_function_endpoints(),
supported_formats: vec![ModelFormat::GGUF],
authentication: vec![AuthType::None, AuthType::BearerToken],
tags: vec![
"llama.cpp".to_string(),
"local".to_string(),
"gguf".to_string(),
"high-performance".to_string(),
],
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config() {
let config = LlamaCppProviderConfig::default();
assert_eq!(config.base_url, "http://localhost:8080");
assert!(config.api_key.is_none());
assert_eq!(config.timeout_secs, 10);
assert!(config.verify_ssl);
assert_eq!(config.health_check_endpoint, "/health"); }
#[test]
fn test_health_response_ok_status() {
let body: serde_json::Value = serde_json::from_str(r#"{"status":"ok"}"#).unwrap();
let status = body.get("status").and_then(|s| s.as_str()).unwrap_or("");
assert_eq!(status, "ok");
}
#[test]
fn test_health_response_loading_status() {
let body: serde_json::Value =
serde_json::from_str(r#"{"status":"loading model"}"#).unwrap();
let status = body.get("status").and_then(|s| s.as_str()).unwrap_or("");
assert_ne!(status, "ok");
}
#[test]
fn test_health_response_missing_status() {
let body: serde_json::Value = serde_json::from_str(r#"{}"#).unwrap();
let status = body.get("status").and_then(|s| s.as_str()).unwrap_or("");
assert_ne!(status, "ok");
}
#[test]
fn test_provider_metadata() {
let meta = LlamaCppProvider::metadata();
assert_eq!(meta.name, "llama.cpp");
assert!(meta.supported_api_types.contains(&ApiType::OpenAI));
assert!(meta.supported_api_types.contains(&ApiType::Anthropic));
assert!(
meta.default_function_endpoints
.contains_key(&ModelFunction::Chat)
);
}
#[test]
fn test_provider_constructs_from_json() {
let cfg = serde_json::json!({
"base_url": "http://example.com:9000",
"timeout_secs": 30
});
let provider =
LlamaCppProvider::new("my-llamacpp", &cfg, &crate::config::Config::default());
assert_eq!(provider.config.base_url, "http://example.com:9000");
assert_eq!(provider.config.timeout_secs, 30);
}
#[test]
fn test_can_run_model_accepts_gguf() {
let provider = LlamaCppProvider::new(
"my-llamacpp",
&serde_json::json!({}),
&crate::config::Config::default(),
);
assert!(provider.can_run_model("gguf", "Q4_K_M"));
assert!(provider.can_run_model("GGUF", "fp16"));
}
#[test]
fn test_can_run_model_rejects_non_gguf() {
let provider = LlamaCppProvider::new(
"my-llamacpp",
&serde_json::json!({}),
&crate::config::Config::default(),
);
assert!(!provider.can_run_model("safetensors", "fp16"));
assert!(!provider.can_run_model("onnx", "fp32"));
}
#[test]
fn test_model_alias_returns_hf_ref_for_gguf_variant() {
let provider = LlamaCppProvider::new(
"my-llamacpp",
&serde_json::json!({}),
&crate::config::Config::default(),
);
let variant = ModelVariant {
format: "GGUF".to_string(),
precision: "Q4_K_M".to_string(),
size_gb: Some(5.3),
url: "https://huggingface.co/ibm-granite/granite-4.1-8b-GGUF/blob/main/granite-4.1-8b-Q4_K_M.gguf".to_string(),
};
assert_eq!(
provider.model_alias("unused".to_string(), Some(&variant)),
Some("ibm-granite/granite-4.1-8b-GGUF:Q4_K_M".to_string())
);
}
#[test]
fn test_model_alias_returns_none_for_ollama_url() {
let provider = LlamaCppProvider::new(
"my-llamacpp",
&serde_json::json!({}),
&crate::config::Config::default(),
);
let variant = ModelVariant {
format: "Ollama".to_string(),
precision: "Q4_K_M".to_string(),
size_gb: Some(5.3),
url: "https://ollama.com/library/granite4.1:8b".to_string(),
};
assert_eq!(
provider.model_alias("unused".to_string(), Some(&variant)),
None
);
}
#[test]
fn test_model_alias_returns_none_when_no_variant() {
let provider = LlamaCppProvider::new(
"my-llamacpp",
&serde_json::json!({}),
&crate::config::Config::default(),
);
assert_eq!(provider.model_alias("unused".to_string(), None), None);
}
#[test]
fn test_find_double_newline() {
assert_eq!(find_double_newline(b"data: {}\n\nmore"), Some(8));
assert_eq!(find_double_newline(b"no terminator here"), None);
}
#[test]
fn test_sse_frame_parses_download_progress() {
let json = r#"{"model":"owner/repo:Q4_K_M","event":"download_progress","data":{"progress":{"https://x/a.gguf":{"done":50,"total":100}}}}"#;
let frame: LlamaCppSseFrame = serde_json::from_str(json).unwrap();
assert_eq!(frame.model, "owner/repo:Q4_K_M");
assert_eq!(frame.event, "download_progress");
let progress = frame.data.get("progress").unwrap().as_object().unwrap();
let entry = progress.values().next().unwrap();
assert_eq!(entry.get("done").unwrap().as_u64(), Some(50));
assert_eq!(entry.get("total").unwrap().as_u64(), Some(100));
}
#[test]
fn test_sse_frame_parses_terminal_events() {
let json = r#"{"model":"owner/repo:Q4_K_M","event":"download_finished","data":{}}"#;
let frame: LlamaCppSseFrame = serde_json::from_str(json).unwrap();
assert_eq!(frame.event, "download_finished");
}
#[test]
fn test_models_response_parses_status() {
let json = r#"{"data":[{"id":"owner/repo:Q4_K_M","status":{"value":"downloading"}}],"object":"list"}"#;
let parsed: LlamaCppModelsResponse = serde_json::from_str(json).unwrap();
assert_eq!(parsed.data[0].id, "owner/repo:Q4_K_M");
assert_eq!(parsed.data[0].status.value, "downloading");
assert!(!parsed.data[0].status.failed);
}
#[test]
fn test_models_response_parses_failed_status() {
let json = r#"{"data":[{"id":"owner/repo:Q4_K_M","status":{"value":"error","failed":true}}],"object":"list"}"#;
let parsed: LlamaCppModelsResponse = serde_json::from_str(json).unwrap();
assert!(parsed.data[0].status.failed);
}
}