docling_rag/embed/
ollama.rs1use super::Embedder;
5use crate::{RagConfig, RagError, Result};
6use async_trait::async_trait;
7use serde::{Deserialize, Serialize};
8
9#[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 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 .error_for_status()?;
61 let body: EmbedResp = resp.json().await?;
62 if body.embeddings.len() != texts.len() {
63 return Err(RagError::Embedding(format!(
64 "ollama returned {} embeddings for {} inputs",
65 body.embeddings.len(),
66 texts.len()
67 )));
68 }
69 Ok(body.embeddings)
70 }
71
72 fn dim(&self) -> usize {
73 self.dim
74 }
75
76 fn id(&self) -> &str {
77 &self.id
78 }
79}