omgbase_search/
provider.rs1use crate::embed::sha256;
5use crate::error::Result;
6
7pub trait EmbeddingProvider {
10 fn model(&self) -> &str;
12 fn dim(&self) -> usize;
14 fn max_input_tokens(&self) -> Option<u32>;
16 fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>>;
18
19 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
28pub const FIXTURE_MODEL: &str = "fixture-hash-8";
30pub const FIXTURE_DIM: usize = 8;
32pub const FIXTURE_MAX_INPUT_TOKENS: u32 = 64;
34
35#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
38pub struct FixtureEmbedder;
39
40impl FixtureEmbedder {
41 #[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 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}