klieo_embed_common/
lib.rs1#![deny(missing_docs)]
2#![deny(rust_2018_idioms)]
3
4use async_trait::async_trait;
20use klieo_core::error::MemoryError;
21
22#[async_trait]
28pub trait Embedder: Send + Sync {
29 fn dimension(&self) -> usize;
33
34 async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, MemoryError>;
36}
37
38pub struct DummyEmbedder;
43
44#[async_trait]
45impl Embedder for DummyEmbedder {
46 fn dimension(&self) -> usize {
47 384
48 }
49
50 async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, MemoryError> {
51 Ok(texts.iter().map(|_| vec![0.0f32; 384]).collect())
52 }
53}
54
55#[cfg(any(test, feature = "test-utils"))]
59pub struct FakeEmbedder {
60 dim: usize,
61}
62
63#[cfg(any(test, feature = "test-utils"))]
64impl FakeEmbedder {
65 pub fn new(dim: usize) -> Self {
67 Self { dim }
68 }
69}
70
71#[cfg(any(test, feature = "test-utils"))]
72impl Default for FakeEmbedder {
73 fn default() -> Self {
74 Self::new(8)
75 }
76}
77
78#[cfg(any(test, feature = "test-utils"))]
79#[async_trait]
80impl Embedder for FakeEmbedder {
81 fn dimension(&self) -> usize {
82 self.dim
83 }
84
85 async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, MemoryError> {
86 let dim = self.dim;
87 Ok(texts
88 .iter()
89 .map(|text| {
90 let mut v = vec![0.0f32; dim];
93 let bytes = text.as_bytes();
94 for (i, slot) in v.iter_mut().enumerate() {
95 const FNV_OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
96 const FNV_PRIME: u64 = 0x0000_0001_0000_01b3;
97 let mut h: u64 = FNV_OFFSET;
98 h ^= i as u64;
99 h = h.wrapping_mul(FNV_PRIME);
100 for &b in bytes {
101 h ^= b as u64;
102 h = h.wrapping_mul(FNV_PRIME);
103 }
104 *slot = (h as f32 / u64::MAX as f32) - 0.5;
105 }
106 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
107 if norm > 0.0 {
108 for x in &mut v {
109 *x /= norm;
110 }
111 }
112 v
113 })
114 .collect())
115 }
116}
117
118#[cfg(test)]
119mod tests {
120 use super::*;
121
122 #[tokio::test]
123 async fn dummy_returns_zero_vectors_of_dim_384() {
124 let e = DummyEmbedder;
125 assert_eq!(e.dimension(), 384);
126 let out = e.embed(&["a".into(), "b".into()]).await.unwrap();
127 assert_eq!(out.len(), 2);
128 assert_eq!(out[0].len(), 384);
129 assert!(out[0].iter().all(|x| *x == 0.0));
130 }
131
132 #[tokio::test]
133 async fn fake_embedder_is_deterministic() {
134 let e = FakeEmbedder::new(16);
135 let a = e.embed(&["hello".into()]).await.unwrap();
136 let b = e.embed(&["hello".into()]).await.unwrap();
137 assert_eq!(a, b);
138 }
139
140 #[tokio::test]
141 async fn fake_embedder_distinguishes_inputs() {
142 let e = FakeEmbedder::new(16);
143 let a = e.embed(&["alpha".into()]).await.unwrap();
144 let b = e.embed(&["beta".into()]).await.unwrap();
145 assert_ne!(a, b);
146 }
147
148 #[tokio::test]
149 async fn fake_embedder_outputs_unit_vectors() {
150 let e = FakeEmbedder::new(8);
151 let v = e.embed(&["hello".into()]).await.unwrap();
152 let norm: f32 = v[0].iter().map(|x| x * x).sum::<f32>().sqrt();
153 assert!(
154 (norm - 1.0).abs() < 1e-5,
155 "fake embedder must produce unit vectors, got norm={norm}"
156 );
157 }
158}