velesdb_memory/
embedder.rs1#[cfg(feature = "ollama")]
10use serde::Deserialize;
11
12#[derive(Debug, thiserror::Error)]
15pub enum EmbedError {
16 #[error("embedding backend error: {0}")]
18 Backend(String),
19 #[error("embedding backend returned an empty vector")]
21 Empty,
22}
23
24pub trait Embedder {
26 fn dimension(&self) -> usize;
28
29 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError>;
34}
35
36#[derive(Debug, Clone)]
42pub struct HashEmbedder {
43 dimension: usize,
44}
45
46impl HashEmbedder {
47 #[must_use]
50 pub fn new(dimension: usize) -> Self {
51 Self { dimension }
52 }
53}
54
55impl Embedder for HashEmbedder {
56 fn dimension(&self) -> usize {
57 self.dimension
58 }
59
60 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError> {
61 let mut vector = vec![0.0_f32; self.dimension];
62 if self.dimension == 0 {
63 return Ok(vector);
64 }
65 let modulus = self.dimension as u64;
66 for token in text.split_whitespace() {
67 let bucket = usize::try_from(crate::id::stable_id(token) % modulus).unwrap_or(0);
68 vector[bucket] += 1.0;
69 }
70 velesdb_core::simd_native::normalize_inplace_native(&mut vector);
71 Ok(vector)
72 }
73}
74
75pub type DynEmbedder = Box<dyn Embedder + Send + Sync>;
79
80impl<T: Embedder + ?Sized> Embedder for Box<T> {
83 fn dimension(&self) -> usize {
84 (**self).dimension()
85 }
86
87 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError> {
88 (**self).embed(text)
89 }
90}
91
92#[cfg(feature = "ollama")]
101pub const DEFAULT_OLLAMA_URL: &str = "http://localhost:11434";
102
103#[cfg(feature = "ollama")]
105pub const DEFAULT_OLLAMA_MODEL: &str = "all-minilm";
106
107#[cfg(feature = "ollama")]
110#[derive(Debug, Clone)]
111pub struct OllamaEmbedder {
112 base_url: String,
113 model: String,
114 dimension: usize,
115}
116
117#[cfg(feature = "ollama")]
118impl OllamaEmbedder {
119 pub fn new(base_url: impl Into<String>, model: impl Into<String>) -> Result<Self, EmbedError> {
126 let base_url = base_url.into();
127 let model = model.into();
128 let dimension = request_embedding(&base_url, &model, "dimension probe")?.len();
129 if dimension == 0 {
130 return Err(EmbedError::Empty);
131 }
132 Ok(Self {
133 base_url,
134 model,
135 dimension,
136 })
137 }
138}
139
140#[cfg(feature = "ollama")]
141impl Embedder for OllamaEmbedder {
142 fn dimension(&self) -> usize {
143 self.dimension
144 }
145
146 fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedError> {
147 request_embedding(&self.base_url, &self.model, text)
148 }
149}
150
151#[cfg(feature = "ollama")]
153fn build_request_body(model: &str, text: &str) -> String {
154 serde_json::json!({ "model": model, "prompt": text }).to_string()
155}
156
157#[cfg(feature = "ollama")]
159#[derive(Deserialize)]
160struct EmbeddingResponse {
161 embedding: Vec<f32>,
162}
163
164#[cfg(feature = "ollama")]
166fn parse_embedding_response(body: &str) -> Result<Vec<f32>, EmbedError> {
167 let parsed: EmbeddingResponse = serde_json::from_str(body)
168 .map_err(|err| EmbedError::Backend(format!("invalid embeddings response: {err}")))?;
169 if parsed.embedding.is_empty() {
170 return Err(EmbedError::Empty);
171 }
172 Ok(parsed.embedding)
173}
174
175#[cfg(feature = "ollama")]
177fn request_embedding(base_url: &str, model: &str, text: &str) -> Result<Vec<f32>, EmbedError> {
178 let url = format!("{base_url}/api/embeddings");
179 let body = build_request_body(model, text);
180 let response = ureq::post(&url)
181 .set("Content-Type", "application/json")
182 .send_string(&body)
183 .map_err(|err| EmbedError::Backend(format!("ollama request failed: {err}")))?;
184 let payload = response
185 .into_string()
186 .map_err(|err| EmbedError::Backend(format!("reading ollama response failed: {err}")))?;
187 parse_embedding_response(&payload)
188}
189
190#[cfg(all(test, feature = "ollama"))]
191mod ollama_tests {
192 use super::*;
193
194 #[test]
195 fn request_body_carries_model_and_prompt() {
196 let body = build_request_body("all-minilm", "hello world");
197 let json: serde_json::Value = serde_json::from_str(&body).expect("valid json");
198 assert_eq!(json["model"], "all-minilm");
199 assert_eq!(json["prompt"], "hello world");
200 }
201
202 #[test]
203 fn parses_a_well_formed_embedding() {
204 let vector = parse_embedding_response(r#"{"embedding":[0.1,0.2,0.3]}"#).expect("parse");
205 assert_eq!(vector.len(), 3);
206 assert!((vector[0] - 0.1_f32).abs() < f32::EPSILON);
207 }
208
209 #[test]
210 fn rejects_an_empty_embedding() {
211 let parsed = parse_embedding_response(r#"{"embedding":[]}"#);
212 assert!(matches!(parsed, Err(EmbedError::Empty)));
213 }
214
215 #[test]
216 fn rejects_a_malformed_response() {
217 let parsed = parse_embedding_response(r#"{"oops":true}"#);
218 assert!(matches!(parsed, Err(EmbedError::Backend(_))));
219 }
220
221 #[test]
222 #[ignore = "requires a local Ollama with an embedding model (ollama pull all-minilm)"]
223 fn embeds_through_a_running_ollama() {
224 let embedder = OllamaEmbedder::new(DEFAULT_OLLAMA_URL, DEFAULT_OLLAMA_MODEL)
225 .expect("connect to ollama");
226 let vector = embedder
227 .embed("parking_lot avoids lock poisoning")
228 .expect("embed");
229 assert_eq!(vector.len(), embedder.dimension());
230 assert!(vector
231 .iter()
232 .any(|&component| component.abs() > f32::EPSILON));
233 }
234}