use crate::error::AppError; use serde::{Deserialize, Serialize}; use std::env; use std::time::Duration; use tracing::debug; #[derive(Clone, Debug)] pub struct OllamaClient { pub base_url: String, pub coder_model: String, pub reasoning_model: String, pub vision_model: String, pub embed_model: String, client: reqwest::Client, } #[derive(Serialize)] struct GenerateRequest<'a> { model: &'a str, prompt: &'a str, #[serde(skip_serializing_if = "Option::is_none")] system: Option<&'a str>, stream: bool, #[serde(skip_serializing_if = "Option::is_none")] images: Option>, } #[derive(Deserialize)] struct GenerateResponse { response: String, } #[derive(Serialize)] struct EmbeddingRequest<'a> { model: &'a str, prompt: &'a str, } #[derive(Deserialize)] struct EmbeddingResponse { embedding: Vec, } impl OllamaClient { pub fn new_from_env() -> Self { let base_url = env::var("OLLAMA_URL").unwrap_or_else(|_| "http://192.168.1.30:11434".to_string()); let coder_model = env::var("OLLAMA_CODER_MODEL").unwrap_or_else(|_| "qwen2.5-coder:1.5b".to_string()); let reasoning_model = env::var("OLLAMA_REASONING_MODEL").unwrap_or_else(|_| "deepseek-r1:1.5b".to_string()); let vision_model = env::var("OLLAMA_VISION_MODEL").unwrap_or_else(|_| "qwen3-vl:2b".to_string()); let embed_model = env::var("OLLAMA_EMBED_MODEL") .unwrap_or_else(|_| "nomic-embed-text:latest".to_string()); let client = reqwest::Client::builder() .timeout(Duration::from_secs(60)) .build() .unwrap_or_default(); Self { base_url, coder_model, reasoning_model, vision_model, embed_model, client, } } /// Health probe check with a strict 1.5-second connection timeout. pub async fn is_available(&self) -> bool { let probe_url = format!("{}/api/tags", self.base_url.trim_end_matches('/')); let probe_client = reqwest::Client::builder() .timeout(Duration::from_millis(1500)) .build(); let client = match probe_client { Ok(c) => c, Err(_) => return false, }; match client.get(&probe_url).send().await { Ok(res) if res.status().is_success() => { debug!("Ollama host at {} is online and responsive.", self.base_url); true } Ok(res) => { debug!("Ollama host returned status {}", res.status()); false } Err(e) => { debug!("Ollama host probe failed (offline/timeout): {}", e); false } } } pub async fn generate( &self, prompt: &str, model_override: Option<&str>, system: Option<&str>, ) -> Result { let model = model_override.unwrap_or(&self.coder_model); let url = format!("{}/api/generate", self.base_url.trim_end_matches('/')); let body = GenerateRequest { model, prompt, system, stream: false, images: None, }; let res = self .client .post(&url) .json(&body) .send() .await .map_err(|e| AppError::Internal(format!("Ollama connection error: {}", e)))?; if !res.status().is_success() { return Err(AppError::Internal(format!( "Ollama API returned HTTP {}", res.status() ))); } let resp_json: GenerateResponse = res.json().await.map_err(|e| { AppError::Internal(format!("Failed to parse Ollama JSON response: {}", e)) })?; Ok(resp_json.response) } pub async fn generate_vision( &self, prompt: &str, image_base64: &str, ) -> Result { let url = format!("{}/api/generate", self.base_url.trim_end_matches('/')); let body = GenerateRequest { model: &self.vision_model, prompt, system: Some( "You are a vision AI assistant. Describe or convert the image provided to code/text as requested.", ), stream: false, images: Some(vec![image_base64]), }; let res = self .client .post(&url) .json(&body) .send() .await .map_err(|e| AppError::Internal(format!("Ollama Vision error: {}", e)))?; if !res.status().is_success() { return Err(AppError::Internal(format!( "Ollama Vision API returned HTTP {}", res.status() ))); } let resp_json: GenerateResponse = res.json().await.map_err(|e| { AppError::Internal(format!("Failed to parse Ollama Vision response: {}", e)) })?; Ok(resp_json.response) } pub async fn embeddings(&self, text: &str) -> Result, AppError> { let url = format!("{}/api/embeddings", self.base_url.trim_end_matches('/')); let body = EmbeddingRequest { model: &self.embed_model, prompt: text, }; let res = self .client .post(&url) .json(&body) .send() .await .map_err(|e| AppError::Internal(format!("Ollama Embeddings error: {}", e)))?; if !res.status().is_success() { return Err(AppError::Internal(format!( "Ollama Embeddings API returned HTTP {}", res.status() ))); } let resp_json: EmbeddingResponse = res.json().await.map_err(|e| { AppError::Internal(format!("Failed to parse Ollama Embeddings response: {}", e)) })?; Ok(resp_json.embedding) } } #[cfg(test)] mod tests { use super::*; #[test] fn test_ollama_client_new_from_env() { let client = OllamaClient::new_from_env(); assert!(!client.base_url.is_empty()); assert!(!client.coder_model.is_empty()); assert!(!client.reasoning_model.is_empty()); assert!(!client.vision_model.is_empty()); assert!(!client.embed_model.is_empty()); } #[tokio::test] async fn test_ollama_client_invalid_host_is_available() { let client = OllamaClient { base_url: "http://127.0.0.1:59999".to_string(), coder_model: "qwen2.5-coder:3b".to_string(), reasoning_model: "deepseek-r1:1.5b".to_string(), vision_model: "qwen3-vl:2b".to_string(), embed_model: "nomic-embed-text:latest".to_string(), client: reqwest::Client::new(), }; assert!(!client.is_available().await); } } #[tokio::test] async fn test_ollama_client_invalid_api_key() { let mut client = OllamaClient::new_from_env(); client.base_url = "http://invalid-api-key:11434".to_string(); assert!(!client.is_available().await); }