Skip to main content

klieo_embed_common/
lib.rs

1#![deny(missing_docs)]
2#![deny(rust_2018_idioms)]
3#![deny(rustdoc::broken_intra_doc_links)]
4
5//! Shared [`Embedder`] trait + dummy/fake implementations for klieo
6//! memory backends.
7//!
8//! Before W3.A17 the trait was duplicated byte-for-byte across
9//! `klieo-memory-sqlite::embedder` and `klieo-memory-qdrant::embedder`
10//! — any downstream embedder (Ollama, OpenAI, fastembed) had to
11//! `impl Embedder` twice. This crate is the single home; both memory
12//! crates re-export from here to keep their public APIs source-stable.
13//!
14//! # Features
15//!
16//! - **Default** — `Embedder` trait + `DummyEmbedder` (zero vectors).
17//! - **`test-utils`** — adds `FakeEmbedder` (deterministic per-text
18//!   hashing) for downstream test harnesses.
19
20use async_trait::async_trait;
21use klieo_core::error::MemoryError;
22
23/// Compute embeddings for one or more texts.
24///
25/// Output vectors must each be of length [`Embedder::dimension`].
26/// Implementations must be deterministic at the type level — the
27/// dimensionality cannot vary across calls on the same instance.
28#[async_trait]
29pub trait Embedder: Send + Sync {
30    /// Embedding dimensionality. Must be constant for a given
31    /// `Embedder` instance — long-term memory backends reject vectors
32    /// of the wrong length at runtime.
33    fn dimension(&self) -> usize;
34
35    /// Compute one embedding per input text.
36    async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, MemoryError>;
37}
38
39/// Default embedder that returns zero vectors.
40///
41/// Lets long-term memory backends store and retrieve facts but cosine
42/// similarity is always 1.0, so recall is FIFO order, not semantic.
43pub 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/// Test-only embedder that hashes each input text into a deterministic
57/// vector. Identical texts produce identical embeddings, so cosine
58/// recall behaves predictably under test.
59#[cfg(any(test, feature = "test-utils"))]
60pub struct FakeEmbedder {
61    dim: usize,
62}
63
64#[cfg(any(test, feature = "test-utils"))]
65impl FakeEmbedder {
66    /// Build a deterministic embedder of the given dimensionality.
67    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                // Deterministic per-text vector via FNV-1a per slot —
92                // toolchain-stable across rustc/std hasher upgrades.
93                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}