Skip to main content

ctx/embeddings/
ollama.rs

1//! Ollama embedding provider.
2//!
3//! Uses a local (or remote) [Ollama](https://ollama.com) server to generate
4//! embeddings via its `/api/embed` endpoint. This gives high-quality embeddings
5//! that run fully offline, without the fastembed model-download constraints and
6//! without OpenAI's per-call cost.
7//!
8//! Unlike OpenAI/fastembed, the embedding dimension is model-dependent
9//! (`nomic-embed-text` = 768, `mxbai-embed-large` = 1024, `qwen3-embedding:8b`
10//! = 4096, …), so it is probed from the model on construction rather than being
11//! a compile-time constant.
12
13use std::sync::OnceLock;
14use std::time::Duration;
15
16use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION, CONTENT_TYPE};
17use reqwest::{Client, StatusCode};
18use serde::{Deserialize, Serialize};
19use tokio::runtime::Runtime;
20
21use super::{Embedding, EmbeddingProvider};
22use crate::error::{CtxError, Result};
23
24/// Default Ollama host when `OLLAMA_HOST` is unset.
25const DEFAULT_HOST: &str = "http://localhost:11434";
26
27/// Default embedding model when `OLLAMA_EMBED_MODEL` is unset.
28const DEFAULT_MODEL: &str = "nomic-embed-text";
29
30const REQUEST_TIMEOUT_SECS: u64 = 60;
31const CONNECT_TIMEOUT_SECS: u64 = 10;
32const MAX_RETRIES: u32 = 3;
33const RETRY_BASE_DELAY_MS: u64 = 500;
34
35/// Global runtime for the sync API when not already in an async context.
36static GLOBAL_RUNTIME: OnceLock<Runtime> = OnceLock::new();
37
38fn get_or_create_runtime() -> &'static Runtime {
39    GLOBAL_RUNTIME.get_or_init(|| {
40        Runtime::new().expect("Failed to create global tokio runtime for Ollama provider")
41    })
42}
43
44/// Ollama `/api/embed` request body. `input` accepts one or many texts.
45#[derive(Serialize)]
46struct OllamaEmbedRequest<'a> {
47    model: &'a str,
48    input: Vec<&'a str>,
49}
50
51/// Ollama `/api/embed` response body.
52#[derive(Deserialize)]
53struct OllamaEmbedResponse {
54    embeddings: Option<Vec<Vec<f32>>>,
55    error: Option<String>,
56}
57
58/// Ollama embedding provider.
59pub struct OllamaProvider {
60    client: Client,
61    host: String,
62    model: String,
63    /// Embedding dimension, probed from the model at construction.
64    dimension: usize,
65}
66
67/// Resolve the host by precedence `OLLAMA_HOST` env > config > default, and
68/// normalize a bare `host:port` (Ollama's own convention) into a URL.
69fn resolve_host(config_host: Option<&str>) -> String {
70    let raw = std::env::var("OLLAMA_HOST")
71        .ok()
72        .filter(|h| !h.is_empty())
73        .or_else(|| config_host.map(str::to_string));
74    match raw {
75        Some(h) if h.starts_with("http://") || h.starts_with("https://") => h,
76        Some(h) => format!("http://{}", h),
77        None => DEFAULT_HOST.to_string(),
78    }
79}
80
81/// Resolve the model by precedence `OLLAMA_EMBED_MODEL` env > config > default.
82fn resolve_model(config_model: Option<&str>) -> String {
83    std::env::var("OLLAMA_EMBED_MODEL")
84        .ok()
85        .filter(|m| !m.is_empty())
86        .or_else(|| config_model.map(str::to_string))
87        .unwrap_or_else(|| DEFAULT_MODEL.to_string())
88}
89
90impl OllamaProvider {
91    /// Create a provider from the environment (`OLLAMA_HOST`, `OLLAMA_EMBED_MODEL`,
92    /// optional `OLLAMA_API_KEY` bearer token), probing the model's dimension
93    /// synchronously. Use [`OllamaProvider::from_env_async`] from async contexts.
94    pub fn from_env() -> Result<Self> {
95        Self::from_config(None, None)
96    }
97
98    /// Create a provider applying config-file `model`/`host` fallbacks (env vars
99    /// still win), probing the dimension synchronously.
100    pub fn from_config(config_model: Option<&str>, config_host: Option<&str>) -> Result<Self> {
101        let mut provider = Self::new_unprobed(config_model, config_host)?;
102        let probe = provider.request(&["dimension probe"])?;
103        provider.dimension = Self::dimension_from_probe(&provider.model, probe)?;
104        Ok(provider)
105    }
106
107    /// Async constructor for use inside an async runtime (e.g. the MCP server),
108    /// where the synchronous probe would deadlock.
109    pub async fn from_env_async() -> Result<Self> {
110        Self::from_config_async(None, None).await
111    }
112
113    /// Async variant of [`OllamaProvider::from_config`].
114    pub async fn from_config_async(
115        config_model: Option<&str>,
116        config_host: Option<&str>,
117    ) -> Result<Self> {
118        let mut provider = Self::new_unprobed(config_model, config_host)?;
119        let probe = provider.request_async(&["dimension probe"]).await?;
120        provider.dimension = Self::dimension_from_probe(&provider.model, probe)?;
121        Ok(provider)
122    }
123
124    /// Build the client/config without probing the dimension (left as 0).
125    fn new_unprobed(config_model: Option<&str>, config_host: Option<&str>) -> Result<Self> {
126        let model = resolve_model(config_model);
127        let host = resolve_host(config_host);
128
129        let mut headers = HeaderMap::new();
130        headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
131        // Optional bearer token for authenticated / remote Ollama hosts.
132        if let Ok(token) = std::env::var("OLLAMA_API_KEY") {
133            if !token.is_empty() {
134                let value = HeaderValue::from_str(&format!("Bearer {}", token)).map_err(|e| {
135                    CtxError::embedding(format!("Invalid OLLAMA_API_KEY format: {}", e))
136                })?;
137                headers.insert(AUTHORIZATION, value);
138            }
139        }
140
141        let client = Client::builder()
142            .timeout(Duration::from_secs(REQUEST_TIMEOUT_SECS))
143            .connect_timeout(Duration::from_secs(CONNECT_TIMEOUT_SECS))
144            .default_headers(headers)
145            .build()
146            .map_err(|e| CtxError::embedding(format!("Failed to build HTTP client: {}", e)))?;
147
148        Ok(Self {
149            client,
150            host,
151            model,
152            dimension: 0,
153        })
154    }
155
156    fn dimension_from_probe(model: &str, probe: Vec<Embedding>) -> Result<usize> {
157        probe
158            .first()
159            .map(|e| e.vector.len())
160            .filter(|d| *d > 0)
161            .ok_or_else(|| {
162                CtxError::embedding(format!(
163                    "Ollama model '{}' returned no embedding on probe",
164                    model
165                ))
166            })
167    }
168
169    /// Async single-text embedding for use inside an async runtime.
170    pub async fn embed_async(&self, text: &str) -> Result<Embedding> {
171        self.request_async(&[text])
172            .await?
173            .pop()
174            .ok_or_else(|| CtxError::embedding("Empty response"))
175    }
176
177    /// The endpoint URL for embeddings.
178    fn embed_url(&self) -> String {
179        format!("{}/api/embed", self.host.trim_end_matches('/'))
180    }
181
182    /// Synchronous request with retry. Errors (rather than deadlocking) if called
183    /// from within an async runtime.
184    fn request(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
185        if tokio::runtime::Handle::try_current().is_ok() {
186            return Err(CtxError::embedding(
187                "Cannot call sync embed() from async context. Use request_async() instead.",
188            ));
189        }
190        get_or_create_runtime().block_on(self.request_async(texts))
191    }
192
193    /// Async request with retry/backoff for transient failures.
194    pub async fn request_async(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
195        let body = OllamaEmbedRequest {
196            model: &self.model,
197            input: texts.to_vec(),
198        };
199
200        let mut last_error = None;
201        for attempt in 0..MAX_RETRIES {
202            match self.send_request(&body).await {
203                Ok(embeddings) => return Ok(embeddings),
204                Err(e) => {
205                    // Retry transient connection / server errors, not "model not
206                    // found" or malformed input.
207                    let retryable = matches!(&e, CtxError::Embedding(msg)
208                        if msg.contains("server error")
209                            || msg.contains("timed out")
210                            || msg.contains("Connection"));
211                    if retryable && attempt < MAX_RETRIES - 1 {
212                        let delay = RETRY_BASE_DELAY_MS * (1 << attempt);
213                        tokio::time::sleep(Duration::from_millis(delay)).await;
214                        last_error = Some(e);
215                        continue;
216                    }
217                    return Err(e);
218                }
219            }
220        }
221        Err(last_error.unwrap_or_else(|| CtxError::embedding("Max retries exceeded")))
222    }
223
224    /// Send a single `/api/embed` request and map the outcome.
225    async fn send_request(&self, body: &OllamaEmbedRequest<'_>) -> Result<Vec<Embedding>> {
226        let response = self
227            .client
228            .post(self.embed_url())
229            .json(body)
230            .send()
231            .await
232            .map_err(|e| {
233                if e.is_timeout() {
234                    CtxError::embedding(format!("Request timed out: {}", e))
235                } else if e.is_connect() {
236                    CtxError::embedding(format!(
237                        "Connection to Ollama at {} failed: {}. Is `ollama serve` running?",
238                        self.host, e
239                    ))
240                } else {
241                    CtxError::embedding(e.to_string())
242                }
243            })?;
244
245        let status = response.status();
246        match status {
247            StatusCode::OK => {
248                let parsed: OllamaEmbedResponse = response
249                    .json()
250                    .await
251                    .map_err(|e| CtxError::embedding(format!("Failed to parse response: {}", e)))?;
252                self.parse_response(parsed)
253            }
254            StatusCode::NOT_FOUND => {
255                // Model not pulled (Ollama returns 404 with an error body).
256                Err(CtxError::ModelNotFound(format!(
257                    "Ollama model '{}' not found. Pull it with: ollama pull {}",
258                    self.model, self.model
259                )))
260            }
261            s if s.is_server_error() => {
262                let body = response.text().await.unwrap_or_default();
263                Err(CtxError::embedding(format!(
264                    "server error ({}): {}",
265                    status, body
266                )))
267            }
268            _ => {
269                let body = response.text().await.unwrap_or_default();
270                // Prefer a structured {"error": ...} message when present.
271                if let Ok(parsed) = serde_json::from_str::<OllamaEmbedResponse>(&body) {
272                    if let Some(err) = parsed.error {
273                        return Err(Self::classify_error(&self.model, err));
274                    }
275                }
276                Err(CtxError::embedding(format!("HTTP {}: {}", status, body)))
277            }
278        }
279    }
280
281    /// Map an Ollama error string to the most specific `CtxError`.
282    fn classify_error(model: &str, message: String) -> CtxError {
283        if message.contains("not found") || message.contains("try pulling") {
284            CtxError::ModelNotFound(format!(
285                "Ollama model '{}' not found. Pull it with: ollama pull {}",
286                model, model
287            ))
288        } else {
289            CtxError::embedding(message)
290        }
291    }
292
293    /// Parse a successful `/api/embed` body into embeddings.
294    fn parse_response(&self, response: OllamaEmbedResponse) -> Result<Vec<Embedding>> {
295        if let Some(err) = response.error {
296            return Err(Self::classify_error(&self.model, err));
297        }
298        let embeddings = response
299            .embeddings
300            .ok_or_else(|| CtxError::embedding("No embeddings in Ollama response"))?;
301        if embeddings.is_empty() {
302            return Err(CtxError::embedding(
303                "Ollama returned an empty embeddings list",
304            ));
305        }
306        // Once the dimension is known, enforce consistency across responses.
307        if self.dimension != 0 {
308            for vector in &embeddings {
309                if vector.len() != self.dimension {
310                    return Err(CtxError::DimensionMismatch {
311                        expected: self.dimension,
312                        actual: vector.len(),
313                    });
314                }
315            }
316        }
317        Ok(embeddings.into_iter().map(Embedding::new).collect())
318    }
319}
320
321impl EmbeddingProvider for OllamaProvider {
322    fn name(&self) -> &str {
323        "ollama"
324    }
325
326    fn dimension(&self) -> usize {
327        self.dimension
328    }
329
330    fn embed(&self, text: &str) -> Result<Embedding> {
331        self.request(&[text])?
332            .pop()
333            .ok_or_else(|| CtxError::embedding("Empty response"))
334    }
335
336    fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
337        // Ollama accepts an array input directly; chunk defensively for very
338        // large batches to bound request size.
339        const BATCH_SIZE: usize = 64;
340        let mut all = Vec::with_capacity(texts.len());
341        for chunk in texts.chunks(BATCH_SIZE) {
342            all.extend(self.request(chunk)?);
343        }
344        Ok(all)
345    }
346}
347
348#[cfg(test)]
349mod tests {
350    use super::*;
351
352    #[test]
353    fn host_precedence_env_over_config_over_default() {
354        std::env::set_var("OLLAMA_HOST", "localhost:11434");
355        assert_eq!(resolve_host(None), "http://localhost:11434"); // bare host normalized
356        assert_eq!(resolve_host(Some("http://cfg:1")), "http://localhost:11434"); // env wins
357        std::env::remove_var("OLLAMA_HOST");
358        assert_eq!(resolve_host(Some("gpu-box:11434")), "http://gpu-box:11434"); // config used
359        assert_eq!(resolve_host(None), DEFAULT_HOST); // default
360    }
361
362    #[test]
363    fn model_precedence_env_over_config_over_default() {
364        std::env::remove_var("OLLAMA_EMBED_MODEL");
365        assert_eq!(resolve_model(None), DEFAULT_MODEL);
366        assert_eq!(
367            resolve_model(Some("qwen3-embedding:8b")),
368            "qwen3-embedding:8b"
369        ); // config
370        std::env::set_var("OLLAMA_EMBED_MODEL", "mxbai-embed-large");
371        assert_eq!(
372            resolve_model(Some("qwen3-embedding:8b")),
373            "mxbai-embed-large"
374        ); // env wins
375        std::env::remove_var("OLLAMA_EMBED_MODEL");
376    }
377
378    /// Build a provider without a probe so `parse_response`/dimension logic can be
379    /// unit-tested offline.
380    fn offline_provider(dimension: usize) -> OllamaProvider {
381        OllamaProvider {
382            client: Client::new(),
383            host: DEFAULT_HOST.to_string(),
384            model: "test-model".to_string(),
385            dimension,
386        }
387    }
388
389    #[test]
390    fn parse_response_success() {
391        let provider = offline_provider(3);
392        let parsed = OllamaEmbedResponse {
393            embeddings: Some(vec![vec![0.1, 0.2, 0.3], vec![0.4, 0.5, 0.6]]),
394            error: None,
395        };
396        let out = provider.parse_response(parsed).unwrap();
397        assert_eq!(out.len(), 2);
398        assert_eq!(out[0].vector, vec![0.1, 0.2, 0.3]);
399    }
400
401    #[test]
402    fn parse_response_dimension_mismatch() {
403        let provider = offline_provider(3);
404        let parsed = OllamaEmbedResponse {
405            embeddings: Some(vec![vec![0.1, 0.2]]), // wrong dim
406            error: None,
407        };
408        assert!(matches!(
409            provider.parse_response(parsed).unwrap_err(),
410            CtxError::DimensionMismatch {
411                expected: 3,
412                actual: 2
413            }
414        ));
415    }
416
417    #[test]
418    fn parse_response_model_not_found() {
419        let provider = offline_provider(0);
420        let parsed = OllamaEmbedResponse {
421            embeddings: None,
422            error: Some("model \"foo\" not found, try pulling it first".to_string()),
423        };
424        assert!(matches!(
425            provider.parse_response(parsed).unwrap_err(),
426            CtxError::ModelNotFound(_)
427        ));
428    }
429
430    #[test]
431    fn parse_response_empty() {
432        let provider = offline_provider(0);
433        let parsed = OllamaEmbedResponse {
434            embeddings: Some(vec![]),
435            error: None,
436        };
437        assert!(provider.parse_response(parsed).is_err());
438    }
439}