1use asupersync::Cx;
18use frankensearch_core::traits::{Embedder, ModelCategory, SearchFuture, l2_normalize};
19
20const FNV_OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
22
23const FNV_PRIME: u64 = 0x0100_0000_01b3;
25
26const MIN_TOKEN_LEN: usize = 2;
28
29const DEFAULT_DIMENSION: usize = 384;
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq)]
34pub enum HashAlgorithm {
35 FnvModular,
40
41 JLProjection {
47 seed: u64,
49 },
50}
51
52#[derive(Debug, Clone)]
67pub struct HashEmbedder {
68 dimension: usize,
69 algorithm: HashAlgorithm,
70}
71
72impl HashEmbedder {
73 #[must_use]
79 pub fn new(dimension: usize, algorithm: HashAlgorithm) -> Self {
80 assert!(dimension > 0, "dimension must be > 0");
81 Self {
82 dimension,
83 algorithm,
84 }
85 }
86
87 #[must_use]
89 pub fn default_384() -> Self {
90 Self::new(DEFAULT_DIMENSION, HashAlgorithm::FnvModular)
91 }
92
93 #[must_use]
95 pub fn default_256() -> Self {
96 Self::new(256, HashAlgorithm::FnvModular)
97 }
98
99 #[must_use]
101 pub fn jl_384(seed: u64) -> Self {
102 Self::new(DEFAULT_DIMENSION, HashAlgorithm::JLProjection { seed })
103 }
104
105 #[must_use]
107 pub fn embed_sync(&self, text: &str) -> Vec<f32> {
108 let tokens = tokenize(text);
109
110 match self.algorithm {
111 HashAlgorithm::FnvModular => self.embed_fnv_modular(&tokens),
112 HashAlgorithm::JLProjection { seed } => self.embed_jl(&tokens, seed),
113 }
114 }
115
116 fn embed_fnv_modular(&self, tokens: &[&str]) -> Vec<f32> {
118 let mut embedding = vec![0.0_f32; self.dimension];
119
120 for token in tokens {
121 let hash = fnv1a_hash(token.as_bytes());
122 #[allow(clippy::cast_possible_truncation)] let index = (hash as usize) % self.dimension;
124 let sign = if (hash >> 63) == 1 { 1.0 } else { -1.0 };
125 embedding[index] += sign;
126 }
127
128 l2_normalize(&embedding)
129 }
130
131 fn embed_jl(&self, tokens: &[&str], seed: u64) -> Vec<f32> {
137 let mut embedding = vec![0.0_f32; self.dimension];
138
139 for token in tokens {
140 let hash = fnv1a_hash(token.as_bytes());
141 let mut state = (seed ^ hash) | 1;
144
145 for dim in &mut embedding {
146 state ^= state << 13;
148 state ^= state >> 7;
149 state ^= state << 17;
150
151 let sign = if (state & 1) == 0 { 1.0 } else { -1.0 };
152 *dim += sign;
153 }
154 }
155
156 l2_normalize(&embedding)
157 }
158}
159
160impl Embedder for HashEmbedder {
161 fn embed<'a>(&'a self, _cx: &'a Cx, text: &'a str) -> SearchFuture<'a, Vec<f32>> {
162 Box::pin(async move { Ok(self.embed_sync(text)) })
164 }
165
166 fn embed_batch<'a>(
167 &'a self,
168 _cx: &'a Cx,
169 texts: &'a [&'a str],
170 ) -> SearchFuture<'a, Vec<Vec<f32>>> {
171 Box::pin(async move { Ok(texts.iter().map(|t| self.embed_sync(t)).collect()) })
172 }
173
174 fn dimension(&self) -> usize {
175 self.dimension
176 }
177
178 fn id(&self) -> &str {
179 match (self.algorithm, self.dimension) {
181 (HashAlgorithm::FnvModular, 384) => "fnv1a-384",
182 (HashAlgorithm::FnvModular, 256) => "fnv1a-256",
183 (HashAlgorithm::JLProjection { .. }, 384) => "jl-384",
184 (HashAlgorithm::JLProjection { .. }, 256) => "jl-256",
185 (HashAlgorithm::FnvModular, _) => "fnv1a-custom",
186 (HashAlgorithm::JLProjection { .. }, _) => "jl-custom",
187 }
188 }
189
190 fn model_name(&self) -> &str {
191 match self.algorithm {
192 HashAlgorithm::FnvModular => "FNV-1a Hash Embedder",
193 HashAlgorithm::JLProjection { .. } => "JL-Projection Hash Embedder",
194 }
195 }
196
197 fn is_semantic(&self) -> bool {
198 false
199 }
200
201 fn category(&self) -> ModelCategory {
202 ModelCategory::HashEmbedder
203 }
204}
205
206fn fnv1a_hash(bytes: &[u8]) -> u64 {
208 let mut hash = FNV_OFFSET;
209 for &byte in bytes {
210 hash ^= u64::from(byte);
211 hash = hash.wrapping_mul(FNV_PRIME);
212 }
213 hash
214}
215
216fn tokenize(text: &str) -> Vec<&str> {
221 text.split(|c: char| !c.is_alphanumeric())
222 .filter(|token| token.len() >= MIN_TOKEN_LEN)
223 .collect()
224}
225
226#[cfg(test)]
227mod tests {
228 use super::*;
229
230 #[test]
233 fn deterministic_same_input_same_output() {
234 let embedder = HashEmbedder::default_384();
235 let a = embedder.embed_sync("hello world");
236 let b = embedder.embed_sync("hello world");
237 assert_eq!(a, b);
238 }
239
240 #[test]
241 fn deterministic_jl_same_seed_same_output() {
242 let embedder = HashEmbedder::jl_384(42);
243 let a = embedder.embed_sync("hello world");
244 let b = embedder.embed_sync("hello world");
245 assert_eq!(a, b);
246 }
247
248 #[test]
249 fn jl_different_seeds_different_output() {
250 let e1 = HashEmbedder::jl_384(42);
251 let e2 = HashEmbedder::jl_384(99);
252 let a = e1.embed_sync("hello world");
253 let b = e2.embed_sync("hello world");
254 assert_ne!(a, b);
255 }
256
257 #[test]
260 fn output_dimension_384() {
261 let embedder = HashEmbedder::default_384();
262 assert_eq!(embedder.embed_sync("test").len(), 384);
263 }
264
265 #[test]
266 fn output_dimension_256() {
267 let embedder = HashEmbedder::default_256();
268 assert_eq!(embedder.embed_sync("test").len(), 256);
269 }
270
271 #[test]
272 fn output_dimension_custom() {
273 let embedder = HashEmbedder::new(128, HashAlgorithm::FnvModular);
274 assert_eq!(embedder.embed_sync("test").len(), 128);
275 }
276
277 #[test]
280 fn output_is_l2_normalized() {
281 let embedder = HashEmbedder::default_384();
282 let vec = embedder.embed_sync("hello world");
283 let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
284 assert!((norm - 1.0).abs() < 1e-6, "norm = {norm}");
285 }
286
287 #[test]
288 fn jl_output_is_l2_normalized() {
289 let embedder = HashEmbedder::jl_384(42);
290 let vec = embedder.embed_sync("hello world");
291 let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
292 assert!((norm - 1.0).abs() < 1e-6, "norm = {norm}");
293 }
294
295 #[test]
298 fn different_inputs_different_embeddings() {
299 let embedder = HashEmbedder::default_384();
300 let a = embedder.embed_sync("hello world");
301 let b = embedder.embed_sync("goodbye universe");
302 assert_ne!(a, b);
303 }
304
305 #[test]
308 fn empty_string_produces_zero_vector() {
309 let embedder = HashEmbedder::default_384();
310 let vec = embedder.embed_sync("");
311 assert_eq!(vec.len(), 384);
313 assert!(vec.iter().all(|&x| x == 0.0));
314 }
315
316 #[test]
317 fn single_char_tokens_filtered() {
318 let embedder = HashEmbedder::default_384();
319 let vec = embedder.embed_sync("a b c");
321 assert!(vec.iter().all(|&x| x == 0.0));
322 }
323
324 #[test]
325 fn long_input_no_panic() {
326 let embedder = HashEmbedder::default_384();
327 let long_text = "word ".repeat(20_000);
328 let vec = embedder.embed_sync(&long_text);
329 assert_eq!(vec.len(), 384);
330 let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
331 assert!((norm - 1.0).abs() < 1e-6);
332 }
333
334 #[test]
337 fn tokenize_basic() {
338 let tokens = tokenize("hello world");
339 assert_eq!(tokens, vec!["hello", "world"]);
340 }
341
342 #[test]
343 fn tokenize_filters_short() {
344 let tokens = tokenize("a bb ccc");
345 assert_eq!(tokens, vec!["bb", "ccc"]);
346 }
347
348 #[test]
349 fn tokenize_splits_on_punctuation() {
350 let tokens = tokenize("hello-world.test");
351 assert_eq!(tokens, vec!["hello", "world", "test"]);
352 }
353
354 #[test]
355 fn tokenize_preserves_case_for_hashing() {
356 let tokens = tokenize("Hello WORLD");
358 assert_eq!(tokens, vec!["Hello", "WORLD"]);
359 }
360
361 #[test]
364 fn fnv1a_empty_is_offset_basis() {
365 assert_eq!(fnv1a_hash(b""), FNV_OFFSET);
366 }
367
368 #[test]
369 fn fnv1a_deterministic() {
370 let a = fnv1a_hash(b"hello");
371 let b = fnv1a_hash(b"hello");
372 assert_eq!(a, b);
373 }
374
375 #[test]
376 fn fnv1a_different_inputs() {
377 assert_ne!(fnv1a_hash(b"hello"), fnv1a_hash(b"world"));
378 }
379
380 #[test]
383 fn embedder_trait_id() {
384 assert_eq!(HashEmbedder::default_384().id(), "fnv1a-384");
385 assert_eq!(HashEmbedder::default_256().id(), "fnv1a-256");
386 assert_eq!(HashEmbedder::jl_384(42).id(), "jl-384");
387 }
388
389 #[test]
390 fn embedder_trait_not_semantic() {
391 assert!(!HashEmbedder::default_384().is_semantic());
392 }
393
394 #[test]
395 fn embedder_trait_category_hash() {
396 assert_eq!(
397 HashEmbedder::default_384().category(),
398 ModelCategory::HashEmbedder
399 );
400 }
401
402 #[test]
403 fn embedder_trait_dimension() {
404 assert_eq!(HashEmbedder::default_384().dimension(), 384);
405 assert_eq!(HashEmbedder::default_256().dimension(), 256);
406 }
407
408 #[test]
409 fn embed_via_trait() {
410 let embedder = HashEmbedder::default_384();
412 let vec = embedder.embed_sync("test query");
413 assert_eq!(vec.len(), 384);
414 }
415
416 #[test]
417 fn embed_batch_via_sync() {
418 let embedder = HashEmbedder::default_384();
419 let texts = ["hello", "world"];
420 let vecs: Vec<_> = texts.iter().map(|t| embedder.embed_sync(t)).collect();
421 assert_eq!(vecs.len(), 2);
422 assert_eq!(vecs[0].len(), 384);
423 assert_eq!(vecs[1].len(), 384);
424 assert_ne!(vecs[0], vecs[1]);
425 }
426
427 #[test]
428 fn embedder_model_name() {
429 assert_eq!(
430 HashEmbedder::default_384().model_name(),
431 "FNV-1a Hash Embedder"
432 );
433 assert_eq!(
434 HashEmbedder::jl_384(42).model_name(),
435 "JL-Projection Hash Embedder"
436 );
437 }
438
439 #[test]
442 fn case_sensitivity_produces_different_embeddings() {
443 let embedder = HashEmbedder::default_384();
444 let lower = embedder.embed_sync("hello world");
445 let upper = embedder.embed_sync("Hello World");
446 assert_ne!(lower, upper);
448
449 let jl = HashEmbedder::jl_384(42);
451 let jl_lower = jl.embed_sync("hello world");
452 let jl_upper = jl.embed_sync("Hello World");
453 assert_ne!(jl_lower, jl_upper);
454 }
455
456 #[test]
457 fn jl_random_pairs_approximately_orthogonal() {
458 use frankensearch_core::cosine_similarity;
459
460 let embedder = HashEmbedder::jl_384(42);
461 let mut total_sim = 0.0_f32;
462 let n: usize = 100;
463
464 for i in 0..n {
465 let text = format!("random document number {i} with unique content");
466 let other = format!("another document {i} about different topics entirely");
467 let a = embedder.embed_sync(&text);
468 let b = embedder.embed_sync(&other);
469 total_sim += cosine_similarity(&a, &b).abs();
470 }
471
472 let mean_sim = total_sim / 100.0_f32;
473 assert!(
474 mean_sim < 0.3,
475 "mean absolute cosine similarity should be low, got {mean_sim}"
476 );
477 }
478}