Skip to main content

docling_rag/embed/
ollama.rs

1//! Ollama embedding provider (the default). Talks to a local Ollama server's
2//! `/api/embed` endpoint. `bge-m3` yields 1024-dimensional vectors.
3
4use super::Embedder;
5use crate::{RagConfig, RagError, Result};
6use async_trait::async_trait;
7use serde::{Deserialize, Serialize};
8
9/// Embedder backed by an Ollama server.
10#[derive(Debug, Clone)]
11pub struct OllamaEmbedder {
12    client: reqwest::Client,
13    base_url: String,
14    model: String,
15    dim: usize,
16    id: String,
17}
18
19#[derive(Serialize)]
20struct EmbedReq<'a> {
21    model: &'a str,
22    input: &'a [String],
23}
24
25#[derive(Deserialize)]
26struct EmbedResp {
27    #[serde(default)]
28    embeddings: Vec<Vec<f32>>,
29}
30
31impl OllamaEmbedder {
32    /// Build from resolved config (`OLLAMA_BASE_URL`, `RAG_EMBED_MODEL`, `RAG_EMBED_DIM`).
33    pub fn from_config(cfg: &RagConfig) -> Self {
34        OllamaEmbedder {
35            client: reqwest::Client::new(),
36            base_url: cfg.ollama_base_url.trim_end_matches('/').to_string(),
37            model: cfg.embed_model.clone(),
38            dim: cfg.embed_dim,
39            id: format!("ollama:{}", cfg.embed_model),
40        }
41    }
42}
43
44#[async_trait]
45impl Embedder for OllamaEmbedder {
46    async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
47        if texts.is_empty() {
48            return Ok(Vec::new());
49        }
50        let url = format!("{}/api/embed", self.base_url);
51        let resp = self
52            .client
53            .post(&url)
54            .json(&EmbedReq {
55                model: &self.model,
56                input: texts,
57            })
58            .send()
59            .await?;
60        // Surface Ollama's error body — it names the actual problem (e.g.
61        // `model "bge-m3" not found, try pulling it first`, which arrives as
62        // a 404 just like an unknown endpoint would on a pre-0.2.6 server).
63        let status = resp.status();
64        if !status.is_success() {
65            let detail = resp.text().await.unwrap_or_default();
66            let hint = if status == reqwest::StatusCode::NOT_FOUND {
67                format!(
68                    " — if the model is missing run `ollama pull {}`; if the \
69                     endpoint is unknown, update Ollama (>= 0.2.6 for /api/embed)",
70                    self.model
71                )
72            } else {
73                String::new()
74            };
75            return Err(RagError::Embedding(format!(
76                "ollama {url}: HTTP {status}: {}{hint}",
77                detail.trim()
78            )));
79        }
80        let body: EmbedResp = resp.json().await?;
81        if body.embeddings.len() != texts.len() {
82            return Err(RagError::Embedding(format!(
83                "ollama returned {} embeddings for {} inputs",
84                body.embeddings.len(),
85                texts.len()
86            )));
87        }
88        Ok(body.embeddings)
89    }
90
91    fn dim(&self) -> usize {
92        self.dim
93    }
94
95    fn id(&self) -> &str {
96        &self.id
97    }
98}