Skip to main content

omgbase_search/
external.rs

1//! External embedding providers (`spec/search` §5): a spawned command spoken
2//! to over newline-delimited JSON, or an http(s) endpoint (feature `http`).
3
4use 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/// The `embedding.*` repo settings (§5).
17#[derive(Clone, Debug, Default, PartialEq, Eq)]
18pub struct EmbeddingSettings {
19    /// A command to spawn or an http(s) URL; `None` → no provider.
20    pub provider: Option<String>,
21    pub model: Option<String>,
22    pub dim: Option<usize>,
23    pub max_input_tokens: Option<u32>,
24}
25
26/// The metadata a provider reports (the stdio handshake line, the HTTP `GET`).
27#[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/// Resolved `model`/`dim`/`max_input_tokens`: the settings, then the
36/// provider's metadata over them.
37#[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/// Whether the provider setting names an endpoint rather than a command.
70#[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/// §5: the `OMGBASE_EMBEDDER_MODEL` / `_DIM` / `_MAX_TOKENS` variables a
78/// spawned command receives for the settings that are set (laid over the
79/// inherited environment; an unset setting leaves the inherited value).
80#[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
95/// Build the provider the settings name: `None` when no provider is set
96/// (`semantic_unavailable` is the caller's), an error when it cannot be
97/// spawned or fails its handshake (`embedder_failed`).
98///
99/// The box is `Send + Sync`: a host shares **one** provider process between
100/// its query path and its background drain thread (`spec/search` §2.6), so
101/// the provider must cross threads and serialize its own pipe traffic (the
102/// [`StdioProvider`] locks around each request; the http provider holds no
103/// connection state).
104pub 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// ---- stdio -----------------------------------------------------------------------
130
131#[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
150/// §5 stdio: a spawned command; one handshake line, then `{id, texts}` →
151/// `{id, vectors}` per request. The child is killed on drop.
152///
153/// `Send + Sync`: the pipes sit behind a mutex, so one instance serves
154/// several threads with one request in flight at a time — a host's query
155/// path and its drain thread share one child (`spec/search` §2.6).
156pub 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    /// Spawn `settings.provider` (split on whitespace: program, args) with the
186    /// [`embedder_env`] over the inherited environment; read the handshake.
187    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// ---- http ---------------------------------------------------------------------------
291
292/// §5 http: `GET url` → metadata (best-effort), `POST url {texts, model}` →
293/// `{vectors}`.
294#[cfg(feature = "http")]
295#[derive(Debug)]
296pub struct HttpProvider {
297    url: String,
298    identity: Identity,
299}
300
301#[cfg(feature = "http")]
302impl HttpProvider {
303    /// Connect: metadata is best-effort (a failure leaves the configured
304    /// `model`/`dim`); `embed` failures are the real signal.
305    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    /// The providers cross threads (one instance is shared by a host's
380    /// query path and its drain thread); this fails to compile otherwise.
381    #[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        // A shell embedder: handshake, then one `{id, vectors}` per request
467        // (two fixed vectors, so requests carry two texts). It also emits an
468        // unrelated id first to prove the reader waits for its own.
469        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}