Skip to main content

omgbase_search/
provider.rs

1//! The embedding provider seam (`spec/search` §5) and the fixture embedder
2//! (§6).
3
4use crate::embed::sha256;
5use crate::error::Result;
6
7/// An embedding model: an external process or endpoint (see
8/// [`crate::external`]) or, in a runner, the [`FixtureEmbedder`].
9pub trait EmbeddingProvider {
10    /// The model name the cache is keyed by.
11    fn model(&self) -> &str;
12    /// The vector length written to `dim`.
13    fn dim(&self) -> usize;
14    /// The model's input limit in tokens, when it reports one (§2.4).
15    fn max_input_tokens(&self) -> Option<u32>;
16    /// Embed a batch of inputs → one float32 vector per input.
17    fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>>;
18
19    /// Embed a bare query string (no context prefix).
20    fn embed_query(&self, query: &str) -> Result<Vec<f32>> {
21        let mut vectors = self.embed(&[query.to_owned()])?;
22        vectors.pop().ok_or_else(|| {
23            crate::Error::EmbedderFailed("embedder returned no vector for the query".to_owned())
24        })
25    }
26}
27
28/// The fixture embedder's model name.
29pub const FIXTURE_MODEL: &str = "fixture-hash-8";
30/// The fixture embedder's dimension.
31pub const FIXTURE_DIM: usize = 8;
32/// The fixture embedder's input limit.
33pub const FIXTURE_MAX_INPUT_TOKENS: u32 = 64;
34
35/// §6: the deterministic hash embedder every runner implements. It is a
36/// runner device, never a product path.
37#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
38pub struct FixtureEmbedder;
39
40impl FixtureEmbedder {
41    /// §6 for one string: `h = sha256(utf8(s))`; `u = h[2i] × 256 + h[2i+1]`,
42    /// `raw[i] = u / 65535 × 2 − 1` in f64 for `i < 8`; L2-normalized in
43    /// f64; each component stored as float32.
44    #[must_use]
45    pub fn vector(s: &str) -> Vec<f32> {
46        let h = sha256(s);
47        let raw: Vec<f64> = (0..FIXTURE_DIM)
48            .map(|i| {
49                let u = f64::from(h[2 * i]) * 256.0 + f64::from(h[2 * i + 1]);
50                u / 65535.0 * 2.0 - 1.0
51            })
52            .collect();
53        let norm = raw.iter().map(|x| x * x).sum::<f64>().sqrt();
54        raw.iter()
55            .map(|x| if norm == 0.0 { 0.0 } else { (x / norm) as f32 })
56            .collect()
57    }
58}
59
60impl EmbeddingProvider for FixtureEmbedder {
61    fn model(&self) -> &str {
62        FIXTURE_MODEL
63    }
64
65    fn dim(&self) -> usize {
66        FIXTURE_DIM
67    }
68
69    fn max_input_tokens(&self) -> Option<u32> {
70        Some(FIXTURE_MAX_INPUT_TOKENS)
71    }
72
73    fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
74        Ok(texts.iter().map(|t| Self::vector(t)).collect())
75    }
76}
77
78#[cfg(test)]
79mod tests {
80    use super::*;
81
82    #[test]
83    fn fixture_embedder_is_the_readme_construction() {
84        // sha256("") = e3b0c442 98fc1c14 9afbf4c8 996fb924 27ae41e4 649b934c a495991b 7852b855
85        let h = sha256("");
86        let raw: Vec<f64> = (0..8)
87            .map(|i| (f64::from(h[2 * i]) * 256.0 + f64::from(h[2 * i + 1])) / 65535.0 * 2.0 - 1.0)
88            .collect();
89        assert_eq!(raw[0], (0xe3b0 as f64) / 65535.0 * 2.0 - 1.0);
90        let norm = raw.iter().map(|x| x * x).sum::<f64>().sqrt();
91        let want: Vec<f32> = raw.iter().map(|x| (x / norm) as f32).collect();
92        assert_eq!(FixtureEmbedder::vector(""), want);
93        let v = FixtureEmbedder::vector("hello");
94        assert_eq!(v.len(), 8);
95        let n: f64 = v.iter().map(|x| f64::from(*x) * f64::from(*x)).sum();
96        assert!((n.sqrt() - 1.0).abs() < 1e-6);
97        assert_ne!(FixtureEmbedder::vector("a"), FixtureEmbedder::vector("b"));
98        let p = FixtureEmbedder;
99        assert_eq!(
100            (p.model(), p.dim(), p.max_input_tokens()),
101            ("fixture-hash-8", 8, Some(64))
102        );
103        assert_eq!(p.embed(&["a".to_owned(), "b".to_owned()]).unwrap().len(), 2);
104        assert_eq!(p.embed_query("a").unwrap(), FixtureEmbedder::vector("a"));
105    }
106}