1use std::collections::HashMap;
5use std::io::{BufRead, BufReader, Write};
6use std::process::{Child, ChildStdin, ChildStdout, Command, Stdio};
7use std::sync::Mutex;
8
9use serde::{Deserialize, Serialize};
10
11use crate::error::{Error, Result};
12use crate::fts::split_ws;
13use crate::provider::EmbeddingProvider;
14use crate::vec::to_f32;
15
16#[derive(Clone, Debug, Default, PartialEq, Eq)]
18pub struct EmbeddingSettings {
19 pub provider: Option<String>,
21 pub model: Option<String>,
22 pub dim: Option<usize>,
23 pub max_input_tokens: Option<u32>,
24}
25
26#[derive(Clone, Debug, Default, Deserialize, PartialEq)]
28struct Metadata {
29 model: Option<String>,
30 dim: Option<f64>,
31 #[serde(rename = "maxInputTokens")]
32 max_input_tokens: Option<f64>,
33}
34
35#[derive(Clone, Debug, PartialEq, Eq)]
38struct Identity {
39 model: String,
40 dim: usize,
41 max_input_tokens: Option<u32>,
42}
43
44impl Identity {
45 fn from_settings(settings: &EmbeddingSettings, default_model: &str) -> Self {
46 Self {
47 model: settings
48 .model
49 .clone()
50 .unwrap_or_else(|| default_model.to_owned()),
51 dim: settings.dim.unwrap_or(0),
52 max_input_tokens: settings.max_input_tokens,
53 }
54 }
55
56 fn apply(&mut self, meta: &Metadata) {
57 if let Some(m) = meta.model.as_deref().filter(|m| !m.is_empty()) {
58 self.model = m.to_owned();
59 }
60 if let Some(d) = meta.dim {
61 self.dim = d as usize;
62 }
63 if let Some(t) = meta.max_input_tokens {
64 self.max_input_tokens = Some(t as u32);
65 }
66 }
67}
68
69#[must_use]
71pub fn is_url(s: &str) -> bool {
72 let t = s.trim();
73 let lower: String = t.chars().take(8).collect::<String>().to_ascii_lowercase();
74 lower.starts_with("http://") || lower.starts_with("https://")
75}
76
77#[must_use]
81pub fn embedder_env(settings: &EmbeddingSettings) -> HashMap<String, String> {
82 let mut env = HashMap::new();
83 if let Some(m) = &settings.model {
84 env.insert("OMGBASE_EMBEDDER_MODEL".to_owned(), m.clone());
85 }
86 if let Some(d) = settings.dim {
87 env.insert("OMGBASE_EMBEDDER_DIM".to_owned(), d.to_string());
88 }
89 if let Some(t) = settings.max_input_tokens {
90 env.insert("OMGBASE_EMBEDDER_MAX_TOKENS".to_owned(), t.to_string());
91 }
92 env
93}
94
95pub fn create_external_provider(
105 settings: &EmbeddingSettings,
106) -> Result<Option<Box<dyn EmbeddingProvider + Send + Sync>>> {
107 let Some(spec) = settings
108 .provider
109 .as_deref()
110 .filter(|p| !p.trim().is_empty())
111 else {
112 return Ok(None);
113 };
114 if is_url(spec) {
115 #[cfg(feature = "http")]
116 {
117 return Ok(Some(Box::new(HttpProvider::connect(settings)?)));
118 }
119 #[cfg(not(feature = "http"))]
120 {
121 return Err(Error::EmbedderFailed(format!(
122 "embedding endpoint {spec} needs the `http` feature of omgbase-search"
123 )));
124 }
125 }
126 Ok(Some(Box::new(StdioProvider::spawn(settings)?)))
127}
128
129#[derive(Serialize)]
132struct Request<'a> {
133 id: u64,
134 texts: &'a [String],
135}
136
137#[derive(Deserialize)]
138struct Response {
139 id: Option<u64>,
140 vectors: Option<Vec<Vec<f64>>>,
141 error: Option<String>,
142}
143
144struct Pipes {
145 stdin: ChildStdin,
146 stdout: BufReader<ChildStdout>,
147 next_id: u64,
148}
149
150pub struct StdioProvider {
157 identity: Identity,
158 command: String,
159 child: Mutex<Child>,
160 pipes: Mutex<Pipes>,
161}
162
163impl std::fmt::Debug for StdioProvider {
164 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
165 f.debug_struct("StdioProvider")
166 .field("command", &self.command)
167 .field("model", &self.identity.model)
168 .field("dim", &self.identity.dim)
169 .finish_non_exhaustive()
170 }
171}
172
173fn read_line(reader: &mut BufReader<ChildStdout>, command: &str) -> Result<String> {
174 let mut line = String::new();
175 let n = reader.read_line(&mut line)?;
176 if n == 0 {
177 return Err(Error::EmbedderFailed(format!(
178 "embedding command '{command}' exited early"
179 )));
180 }
181 Ok(line.trim_end_matches(['\n', '\r']).to_owned())
182}
183
184impl StdioProvider {
185 pub fn spawn(settings: &EmbeddingSettings) -> Result<Self> {
188 let spec = settings.provider.as_deref().unwrap_or_default();
189 let mut parts = split_ws(spec);
190 let Some(program) = parts.next() else {
191 return Err(Error::EmbedderFailed(
192 "embedding command is empty".to_owned(),
193 ));
194 };
195 let args: Vec<&str> = parts.collect();
196 let mut child = Command::new(program)
197 .args(&args)
198 .envs(embedder_env(settings))
199 .stdin(Stdio::piped())
200 .stdout(Stdio::piped())
201 .stderr(Stdio::inherit())
202 .spawn()
203 .map_err(|e| {
204 Error::EmbedderFailed(format!("embedding command '{spec}' failed to spawn: {e}"))
205 })?;
206 let stdin = child.stdin.take().expect("piped stdin");
207 let stdout = BufReader::new(child.stdout.take().expect("piped stdout"));
208 let mut pipes = Pipes {
209 stdin,
210 stdout,
211 next_id: 1,
212 };
213 let handshake = match read_line(&mut pipes.stdout, program) {
214 Ok(l) => l,
215 Err(e) => {
216 let _ = child.kill();
217 return Err(e);
218 }
219 };
220 let meta: Metadata = serde_json::from_str(&handshake).map_err(|_| {
221 let _ = child.kill();
222 Error::EmbedderFailed(format!(
223 "embedding command '{program}' sent an invalid handshake line: {}",
224 handshake.chars().take(120).collect::<String>()
225 ))
226 })?;
227 let mut identity = Identity::from_settings(settings, "stdio");
228 identity.apply(&meta);
229 Ok(Self {
230 identity,
231 command: program.to_owned(),
232 child: Mutex::new(child),
233 pipes: Mutex::new(pipes),
234 })
235 }
236}
237
238impl Drop for StdioProvider {
239 fn drop(&mut self) {
240 if let Ok(mut child) = self.child.lock() {
241 let _ = child.kill();
242 let _ = child.wait();
243 }
244 }
245}
246
247impl EmbeddingProvider for StdioProvider {
248 fn model(&self) -> &str {
249 &self.identity.model
250 }
251
252 fn dim(&self) -> usize {
253 self.identity.dim
254 }
255
256 fn max_input_tokens(&self) -> Option<u32> {
257 self.identity.max_input_tokens
258 }
259
260 fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
261 if texts.is_empty() {
262 return Ok(Vec::new());
263 }
264 let mut pipes = self
265 .pipes
266 .lock()
267 .map_err(|_| Error::EmbedderFailed("embedder pipes poisoned".to_owned()))?;
268 let id = pipes.next_id;
269 pipes.next_id += 1;
270 let mut line = serde_json::to_string(&Request { id, texts })?;
271 line.push('\n');
272 pipes.stdin.write_all(line.as_bytes())?;
273 pipes.stdin.flush()?;
274 loop {
275 let raw = read_line(&mut pipes.stdout, &self.command)?;
276 let resp: Response = serde_json::from_str(&raw)?;
277 if let Some(err) = resp.error {
278 return Err(Error::EmbedderFailed(format!("embedder error: {err}")));
279 }
280 if resp.id == Some(id) {
281 let vectors = resp.vectors.ok_or_else(|| {
282 Error::EmbedderFailed("embedder response missing vectors".to_owned())
283 })?;
284 return Ok(vectors.iter().map(|v| to_f32(v)).collect());
285 }
286 }
287 }
288}
289
290#[cfg(feature = "http")]
295#[derive(Debug)]
296pub struct HttpProvider {
297 url: String,
298 identity: Identity,
299}
300
301#[cfg(feature = "http")]
302impl HttpProvider {
303 pub fn connect(settings: &EmbeddingSettings) -> Result<Self> {
306 let url = settings
307 .provider
308 .as_deref()
309 .unwrap_or_default()
310 .trim()
311 .to_owned();
312 let mut identity = Identity::from_settings(settings, "http");
313 if let Ok(mut res) = ureq::get(&url).call() {
314 if let Ok(meta) = res.body_mut().read_json::<Metadata>() {
315 identity.apply(&meta);
316 }
317 }
318 Ok(Self { url, identity })
319 }
320}
321
322#[cfg(feature = "http")]
323impl EmbeddingProvider for HttpProvider {
324 fn model(&self) -> &str {
325 &self.identity.model
326 }
327
328 fn dim(&self) -> usize {
329 self.identity.dim
330 }
331
332 fn max_input_tokens(&self) -> Option<u32> {
333 self.identity.max_input_tokens
334 }
335
336 fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
337 if texts.is_empty() {
338 return Ok(Vec::new());
339 }
340 #[derive(Serialize)]
341 struct Body<'a> {
342 texts: &'a [String],
343 model: &'a str,
344 }
345 #[derive(Deserialize)]
346 struct Reply {
347 vectors: Option<Vec<Vec<f64>>>,
348 }
349 let mut res = ureq::post(&self.url)
350 .send_json(Body {
351 texts,
352 model: &self.identity.model,
353 })
354 .map_err(|e| match e {
355 ureq::Error::StatusCode(code) => Error::EmbedderFailed(format!(
356 "embedding endpoint {} returned {code}",
357 self.url
358 )),
359 other => Error::EmbedderFailed(format!("embedding endpoint {}: {other}", self.url)),
360 })?;
361 let reply: Reply = res
362 .body_mut()
363 .read_json()
364 .map_err(|e| Error::EmbedderFailed(format!("embedding endpoint {}: {e}", self.url)))?;
365 let vectors = reply.vectors.ok_or_else(|| {
366 Error::EmbedderFailed(format!(
367 "embedding endpoint {} returned no \"vectors\"",
368 self.url
369 ))
370 })?;
371 Ok(vectors.iter().map(|v| to_f32(v)).collect())
372 }
373}
374
375#[cfg(test)]
376mod tests {
377 use super::*;
378
379 #[test]
382 fn providers_are_send_and_sync() {
383 fn assert_shareable<T: Send + Sync>() {}
384 assert_shareable::<StdioProvider>();
385 #[cfg(feature = "http")]
386 assert_shareable::<HttpProvider>();
387 assert_shareable::<Box<dyn EmbeddingProvider + Send + Sync>>();
388 }
389
390 #[test]
391 fn settings_env_and_url_detection() {
392 assert!(is_url("https://x.example/embed"));
393 assert!(is_url(" HTTP://x "));
394 assert!(!is_url("omgbase-embedder"));
395 assert!(!is_url("httpd --serve"));
396 let s = EmbeddingSettings {
397 provider: Some("cmd".to_owned()),
398 model: Some("m".to_owned()),
399 dim: Some(3),
400 max_input_tokens: None,
401 };
402 let env = embedder_env(&s);
403 assert_eq!(
404 env.get("OMGBASE_EMBEDDER_MODEL").map(String::as_str),
405 Some("m")
406 );
407 assert_eq!(
408 env.get("OMGBASE_EMBEDDER_DIM").map(String::as_str),
409 Some("3")
410 );
411 assert!(!env.contains_key("OMGBASE_EMBEDDER_MAX_TOKENS"));
412 assert!(
413 create_external_provider(&EmbeddingSettings::default())
414 .unwrap()
415 .is_none()
416 );
417 }
418
419 #[test]
420 fn identity_precedence() {
421 let s = EmbeddingSettings {
422 provider: None,
423 model: Some("cfg".to_owned()),
424 dim: Some(2),
425 max_input_tokens: Some(10),
426 };
427 let mut id = Identity::from_settings(&s, "stdio");
428 id.apply(&Metadata {
429 model: Some("hand".to_owned()),
430 dim: Some(4.0),
431 max_input_tokens: None,
432 });
433 assert_eq!(
434 id,
435 Identity {
436 model: "hand".to_owned(),
437 dim: 4,
438 max_input_tokens: Some(10)
439 }
440 );
441 let id = Identity::from_settings(&EmbeddingSettings::default(), "http");
442 assert_eq!(
443 id,
444 Identity {
445 model: "http".to_owned(),
446 dim: 0,
447 max_input_tokens: None
448 }
449 );
450 }
451
452 #[test]
453 fn spawn_failure_is_embedder_failed() {
454 let s = EmbeddingSettings {
455 provider: Some("/definitely/not/a/program".to_owned()),
456 ..EmbeddingSettings::default()
457 };
458 let err = create_external_provider(&s).err().expect("spawn fails");
459 assert_eq!(err.code(), "embedder_failed");
460 assert!(err.to_string().contains("failed to spawn"));
461 }
462
463 #[cfg(unix)]
464 #[test]
465 fn stdio_protocol_round_trip() {
466 let script = r#"echo '{"model":"sh","dim":2,"maxInputTokens":32}'
470while IFS= read -r line; do
471 id=$(printf '%s' "$line" | sed -E 's/.*"id":([0-9]+).*/\1/')
472 echo '{"id":999,"vectors":[[9,9]]}'
473 echo "{\"id\":$id,\"vectors\":[[1,0.5],[0,1]]}"
474done"#;
475 let dir = std::env::temp_dir().join(format!("omgbase-search-{}", std::process::id()));
476 std::fs::create_dir_all(&dir).unwrap();
477 let path = dir.join("embedder.sh");
478 std::fs::write(&path, script).unwrap();
479 let settings = EmbeddingSettings {
480 provider: Some(format!("sh {}", path.display())),
481 model: Some("cfg".to_owned()),
482 dim: Some(9),
483 max_input_tokens: None,
484 };
485 let p = StdioProvider::spawn(&settings).expect("spawns");
486 assert_eq!(
487 (p.model(), p.dim(), p.max_input_tokens()),
488 ("sh", 2, Some(32))
489 );
490 let v = p.embed(&["a".to_owned(), "b".to_owned()]).unwrap();
491 assert_eq!(v, vec![vec![1.0f32, 0.5], vec![0.0, 1.0]]);
492 let v = p.embed(&["c".to_owned(), "d".to_owned()]).unwrap();
493 assert_eq!(v.len(), 2);
494 assert_eq!(p.embed(&[]).unwrap(), Vec::<Vec<f32>>::new());
495 drop(p);
496 let _ = std::fs::remove_dir_all(&dir);
497 }
498}