Skip to main content

rig_core/providers/ollama/
embedding.rs

1//! Ollama's embedding wire, `POST /api/embed`, and its embedding models.
2//!
3//! ```
4//! use rig_core::providers::ollama::{ALL_MINILM, OllamaConfig};
5//! let model = OllamaConfig::new().client().embedding(ALL_MINILM, None);
6//! assert_eq!(model.wire.ndims, 384);
7//! ```
8
9use crate::error::{EncodeError, ProviderError};
10use crate::operation::Embedding;
11use crate::wire::{
12    Body, Capabilities, Decoder, Descriptor, Encoded, Flow, Framing, Mode, Out, Wire, WireEvent,
13    WireFrame,
14};
15use serde::{Deserialize, Serialize};
16
17use super::{OllamaConfig, PROVIDER_NAME};
18
19/// The `all-minilm` embedding model.
20pub const ALL_MINILM: &str = "all-minilm";
21/// The `nomic-embed-text` embedding model.
22pub const NOMIC_EMBED_TEXT: &str = "nomic-embed-text";
23/// The `mxbai-embed-large` embedding model.
24pub const MXBAI_EMBED_LARGE: &str = "mxbai-embed-large";
25/// The `bge-m3` multilingual embedding model.
26pub const BGE_M3: &str = "bge-m3";
27/// The `embeddinggemma` embedding model.
28pub const EMBEDDINGGEMMA: &str = "embeddinggemma";
29/// The `qwen3-embedding` embedding model family; dimensions vary by size, so pass them explicitly.
30pub const QWEN3_EMBEDDING: &str = "qwen3-embedding";
31
32fn model_dimensions_from_identifier(identifier: &str) -> Option<usize> {
33    match identifier {
34        ALL_MINILM => Some(384),
35        NOMIC_EMBED_TEXT => Some(768),
36        MXBAI_EMBED_LARGE => Some(1024),
37        BGE_M3 => Some(1024),
38        EMBEDDINGGEMMA => Some(768),
39        _ => None,
40    }
41}
42
43#[derive(Debug, Clone, Serialize, Deserialize)]
44pub struct EmbeddingResponse {
45    pub model: String,
46    pub embeddings: Vec<Vec<f64>>,
47    #[serde(default)]
48    pub total_duration: Option<u64>,
49    #[serde(default)]
50    pub load_duration: Option<u64>,
51    #[serde(default)]
52    pub prompt_eval_count: Option<u64>,
53}
54
55/// The most texts `POST /api/embed` accepts in one call.
56const MAX_DOCUMENTS: usize = 1024;
57
58impl OllamaConfig {
59    /// Build an embedding wire reporting the supplied width, known model width,
60    /// or zero if unknown. The width is metadata and is not sent to the daemon.
61    pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
62        let model = model.into();
63        let ndims = ndims
64            .or_else(|| model_dimensions_from_identifier(&model))
65            .unwrap_or_default();
66        Embeddings {
67            provider: self.clone(),
68            model,
69            ndims,
70        }
71    }
72}
73
74/// The embedding wire: `POST /api/embed`.
75#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
76pub struct Embeddings {
77    /// The daemon this wire speaks to.
78    pub provider: OllamaConfig,
79    /// The model to address.
80    pub model: String,
81    /// The width this wire reports, from the caller or the model's published
82    /// dimensions. `0` means neither named one.
83    pub ndims: usize,
84}
85
86impl Wire for Embeddings {
87    type Op = Embedding;
88    type Payload = crate::wire::Encoded;
89    type Frame = crate::wire::WireFrame;
90    type Decoder<'id> = EmbeddingsDecoder;
91    type Reassembler = crate::wire::document::Unreassembled;
92
93    fn describe(&self) -> Descriptor<'_> {
94        Descriptor::new(PROVIDER_NAME)
95            .model(self.model.as_str())
96            .capabilities(Capabilities::embedding(MAX_DOCUMENTS, self.ndims))
97    }
98
99    fn encode(&self, texts: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
100        let body = serde_json::json!({ "model": self.model, "input": texts });
101        let request = self
102            .provider
103            .request(http::Method::POST, "/api/embed")
104            .body(Body::Bytes(serde_json::to_vec(&body)?))?;
105        Ok(Encoded::new(request, Framing::Whole))
106    }
107
108    fn decoder<'id>(&self) -> Self::Decoder<'id> {
109        EmbeddingsDecoder
110    }
111}
112
113/// Decodes one `/api/embed` reply.
114pub struct EmbeddingsDecoder;
115
116impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
117    type Event = EmbeddingResponse;
118
119    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
120        crate::providers::internal::wire::classify_marker_keyed_frame(
121            &frame.as_str(),
122            &["embeddings"],
123        )
124    }
125
126    fn decode(
127        &mut self,
128        reply: Self::Event,
129        out: Out<'id, Embedding>,
130    ) -> Result<Flow, ProviderError> {
131        // Ollama counts the prompt it embedded and nothing else: every token
132        // of an embedding is input.
133        let usage = crate::completion::Usage {
134            input_tokens: reply.prompt_eval_count,
135            total_tokens: reply.prompt_eval_count,
136            ..Default::default()
137        };
138        Ok(out.end(crate::embeddings::EmbeddingResponse {
139            model: Some(reply.model),
140            usage,
141            ..crate::embeddings::EmbeddingResponse::from_vectors(reply.embeddings)
142        }))
143    }
144}
145
146#[cfg(test)]
147mod tests;