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