1use crate::error::HostError;
11
12pub trait Embedder: Send + Sync {
30 fn dim(&self) -> usize;
33
34 fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError>;
44}
45
46#[derive(Clone, Copy, Debug, Default)]
49pub struct NullEmbedder;
50
51impl Embedder for NullEmbedder {
52 fn dim(&self) -> usize {
53 0
54 }
55
56 fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
57 Ok(vec![Vec::new(); texts.len()])
58 }
59}
60
61#[derive(Clone)]
73pub struct SharedEmbedder(std::sync::Arc<dyn Embedder>);
74
75impl SharedEmbedder {
76 pub fn new(inner: Box<dyn Embedder>) -> Self {
78 Self(std::sync::Arc::from(inner))
79 }
80}
81
82impl std::fmt::Debug for SharedEmbedder {
83 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
84 f.debug_struct("SharedEmbedder")
85 .field("dim", &self.0.dim())
86 .finish()
87 }
88}
89
90impl Embedder for SharedEmbedder {
91 fn dim(&self) -> usize {
92 self.0.dim()
93 }
94
95 fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
96 self.0.embed(texts)
97 }
98}
99
100#[derive(Debug)]
102pub struct OpenAiCompatEmbedder {
103 url: String,
104 model: String,
105 api_key: Option<String>,
106 dim: usize,
107 agent: ureq::Agent,
108}
109
110impl OpenAiCompatEmbedder {
111 pub fn new(base_url: &str, model: &str, dim: usize) -> Self {
117 Self {
118 url: format!("{}/embeddings", base_url.trim_end_matches('/')),
119 model: model.to_string(),
120 api_key: None,
121 dim,
122 agent: ureq::Agent::new_with_defaults(),
123 }
124 }
125
126 pub fn with_api_key(mut self, key: impl Into<String>) -> Self {
129 self.api_key = Some(key.into());
130 self
131 }
132}
133
134impl Embedder for OpenAiCompatEmbedder {
135 fn dim(&self) -> usize {
136 self.dim
137 }
138
139 fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
140 if texts.is_empty() {
141 return Ok(Vec::new());
142 }
143 let body = serde_json::json!({ "model": self.model, "input": texts });
144 let mut request = self.agent.post(&self.url);
145 if let Some(key) = &self.api_key {
146 request = request.header("Authorization", &format!("Bearer {key}"));
147 }
148 let mut response = request
149 .send_json(&body)
150 .map_err(|e| HostError::Embed(format!("request to {}: {e}", self.url)))?;
151 let value: serde_json::Value = response
152 .body_mut()
153 .read_json()
154 .map_err(|e| HostError::Embed(format!("response body: {e}")))?;
155
156 let data = value
160 .get("data")
161 .and_then(|d| d.as_array())
162 .ok_or_else(|| HostError::Embed("response has no data array".into()))?;
163 if data.len() != texts.len() {
164 return Err(HostError::Embed(format!(
165 "expected {} embeddings, got {}",
166 texts.len(),
167 data.len()
168 )));
169 }
170 let mut out = vec![Vec::new(); texts.len()];
171 for item in data {
172 let index = item
173 .get("index")
174 .and_then(|i| i.as_u64())
175 .ok_or_else(|| HostError::Embed("embedding without an index".into()))?
176 as usize;
177 let raw = item
178 .get("embedding")
179 .and_then(|e| e.as_array())
180 .ok_or_else(|| HostError::Embed("embedding is not an array".into()))?;
181 if index >= out.len() || !out[index].is_empty() {
182 return Err(HostError::Embed(format!("bad embedding index {index}")));
183 }
184 if raw.len() != self.dim {
185 return Err(HostError::Embed(format!(
186 "dimension mismatch: server sent {}, configured {}",
187 raw.len(),
188 self.dim
189 )));
190 }
191 let mut v = Vec::with_capacity(raw.len());
192 for x in raw {
193 v.push(
194 x.as_f64().ok_or_else(|| {
195 HostError::Embed("embedding component is not a number".into())
196 })? as f32,
197 );
198 }
199 out[index] = v;
200 }
201 Ok(out)
202 }
203}
204
205#[cfg(test)]
206mod tests {
207 use super::*;
208
209 use std::sync::atomic::{AtomicUsize, Ordering};
210
211 struct Counting(AtomicUsize);
215 impl Embedder for Counting {
216 fn dim(&self) -> usize {
217 3
218 }
219 fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
220 let total = self.0.fetch_add(texts.len(), Ordering::Relaxed) + texts.len();
221 Ok(vec![vec![total as f32; 3]; texts.len()])
222 }
223 }
224
225 struct Overlapping {
228 inside: AtomicUsize,
229 peak: AtomicUsize,
230 }
231 impl Embedder for Overlapping {
232 fn dim(&self) -> usize {
233 1
234 }
235 fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, HostError> {
236 let now = self.inside.fetch_add(1, Ordering::SeqCst) + 1;
237 self.peak.fetch_max(now, Ordering::SeqCst);
238 std::thread::sleep(std::time::Duration::from_millis(50));
239 self.inside.fetch_sub(1, Ordering::SeqCst);
240 Ok(vec![vec![0.0]; texts.len()])
241 }
242 }
243
244 #[test]
245 fn clones_of_a_shared_embedder_reach_the_same_provider() {
246 let shared = SharedEmbedder::new(Box::new(Counting(AtomicUsize::new(0))));
247 let a = shared.clone();
248 let b = shared.clone();
249
250 assert_eq!(a.dim(), 3);
251 assert_eq!(format!("{shared:?}"), "SharedEmbedder { dim: 3 }");
252
253 assert_eq!(a.embed(&["x"]).unwrap(), vec![vec![1.0; 3]]);
256 assert_eq!(b.embed(&["y", "z"]).unwrap(), vec![vec![3.0; 3]; 2]);
257 }
258
259 #[test]
260 fn concurrent_callers_are_inside_the_provider_at_the_same_time() {
261 let provider = std::sync::Arc::new(Overlapping {
266 inside: AtomicUsize::new(0),
267 peak: AtomicUsize::new(0),
268 });
269 let shared = SharedEmbedder(provider.clone());
270
271 std::thread::scope(|scope| {
272 for _ in 0..4 {
273 let handle = shared.clone();
274 scope.spawn(move || handle.embed(&["question"]).unwrap());
275 }
276 });
277
278 assert!(
279 provider.peak.load(Ordering::SeqCst) > 1,
280 "callers serialized: peak concurrency was {}",
281 provider.peak.load(Ordering::SeqCst)
282 );
283 }
284
285 #[test]
286 fn the_null_embedder_produces_one_empty_vector_per_text() {
287 let null = NullEmbedder;
288 assert_eq!(null.dim(), 0);
289 assert_eq!(null.embed(&["a", "b"]).unwrap(), vec![Vec::<f32>::new(); 2]);
290 }
291}