rig_core/providers/ollama/
embedding.rs1use 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
19pub const ALL_MINILM: &str = "all-minilm";
21pub const NOMIC_EMBED_TEXT: &str = "nomic-embed-text";
23pub const MXBAI_EMBED_LARGE: &str = "mxbai-embed-large";
25pub const BGE_M3: &str = "bge-m3";
27pub const EMBEDDINGGEMMA: &str = "embeddinggemma";
29pub 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
55const MAX_DOCUMENTS: usize = 1024;
57
58impl OllamaConfig {
59 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#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
76pub struct Embeddings {
77 pub provider: OllamaConfig,
79 pub model: String,
81 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
113pub 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 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;