Skip to main content

harn_hostlib/embed/
backend.rs

1//! Pluggable embedding backends.
2//!
3//! The [`Embedder`] trait is the futureproof seam: every backend turns a
4//! string into a fixed-dimension `f32` vector that downstream cosine math
5//! ranks. Two pure-Rust, zero-asset, fully-offline backends ship in-tree:
6//!
7//! * [`LexicalEmbedder`] — the always-available default. A hashed
8//!   bag-of-features (word tokens + char trigrams) projected into a fixed
9//!   dimension via the hashing trick, then L2-normalized. No model, no
10//!   asset, no network; microsecond latency; deterministic across OSes.
11//!   This is the graceful-degradation floor: it is what every other
12//!   backend falls back to when its asset is missing.
13//!
14//! * [`StaticEmbedder`] — a Model2Vec / "potion"-style static
15//!   token-pooled embedder. Loads precomputed per-token vectors from a
16//!   resolved-on-disk asset, then `tokenize -> lookup -> mean -> normalize`
17//!   with no neural-network inference. Microsecond latency, ~92% of
18//!   MiniLM-class quality when the asset is present. Constructed via
19//!   [`StaticEmbedder::from_asset_dir`], which fails cleanly (so callers
20//!   fall back to lexical) when the asset is absent.
21//!
22//! A higher-accuracy on-device transformer backend (candle / ONNX) can be
23//! added later behind a Cargo feature without changing this trait or any
24//! consumer: implement [`Embedder`], resolve its asset the same way, and
25//! register it as the active backend when the feature + setting are on.
26
27use std::collections::HashMap;
28use std::path::{Path, PathBuf};
29
30pub use harn_session_store::{Embedder, LexicalEmbedder};
31
32use super::tokenize;
33
34/// Model2Vec / potion-style static token-pooled embedder.
35///
36/// Holds a precomputed `token -> vector` table loaded from a vendored
37/// asset. Embedding is `tokenize -> lookup each token's vector -> mean ->
38/// L2-normalize`, with NO neural-network forward pass — that is the entire
39/// point of static embeddings (microsecond latency, tiny footprint).
40///
41/// ## Asset format (intentionally simple + dependency-free)
42///
43/// The asset directory contains `static-embeddings.json`:
44/// ```json
45/// { "dim": 8, "vectors": { "rate": [...8 floats...], "limit": [...] } }
46/// ```
47/// Tokens are the same word tokens [`tokenize::word_tokens`] produces, so a
48/// distilled potion table can be exported into this shape offline. A real
49/// `.safetensors` loader (via `model2vec-rs`) can be added behind a feature
50/// later without touching this trait or the JSON fallback — both satisfy
51/// the same `token -> vector` contract.
52pub struct StaticEmbedder {
53    dim: usize,
54    vectors: HashMap<String, Vec<f32>>,
55    name: String,
56    /// Lexical fallback used when a query contains *no* known tokens, so a
57    /// previously-unseen identifier still gets a non-degenerate vector
58    /// instead of collapsing to zero.
59    fallback: LexicalEmbedder,
60}
61
62impl StaticEmbedder {
63    /// Resolve and load `static-embeddings.json` under `asset_dir`.
64    ///
65    /// Returns `Err` (so the caller can fall back to lexical) when the
66    /// directory or file is missing, unreadable, malformed, or empty. This
67    /// is the sandbox/settings-aware degradation contract: a missing asset
68    /// never panics and never blocks — it just selects the lexical floor.
69    pub fn from_asset_dir(asset_dir: &Path) -> Result<Self, String> {
70        let path = asset_dir.join("static-embeddings.json");
71        let raw = std::fs::read_to_string(&path)
72            .map_err(|e| format!("static embedding asset {} unreadable: {e}", path.display()))?;
73        Self::from_json(&raw)
74    }
75
76    /// Parse an in-memory asset document. Split out for testability and so
77    /// future loaders (safetensors) can reuse the validation.
78    pub fn from_json(raw: &str) -> Result<Self, String> {
79        // Hand-rolled minimal parse keeps the default build dependency-free
80        // (no serde_json pulled in just for an optional asset). The format
81        // is small and fixed; we accept the documented shape only.
82        let doc: AssetDoc = parse_asset(raw)?;
83        if doc.vectors.is_empty() {
84            return Err("static embedding asset has no vectors".to_string());
85        }
86        for (tok, v) in &doc.vectors {
87            if v.len() != doc.dim {
88                return Err(format!(
89                    "static embedding vector for `{tok}` has length {} but dim is {}",
90                    v.len(),
91                    doc.dim
92                ));
93            }
94        }
95        Ok(Self {
96            dim: doc.dim,
97            vectors: doc.vectors,
98            name: "static-model2vec".to_string(),
99            fallback: LexicalEmbedder::new(doc.dim),
100        })
101    }
102}
103
104impl Embedder for StaticEmbedder {
105    fn embed(&self, text: &str) -> Vec<f32> {
106        let mut acc = vec![0.0f32; self.dim];
107        let mut hits = 0usize;
108        for token in tokenize::word_tokens(text) {
109            if let Some(v) = self.vectors.get(&token) {
110                for (a, x) in acc.iter_mut().zip(v.iter()) {
111                    *a += x;
112                }
113                hits += 1;
114            }
115        }
116        if hits == 0 {
117            // No known tokens: fall back to the lexical projection so a
118            // novel identifier is still comparable rather than all-zero.
119            return self.fallback.embed(text);
120        }
121        let inv = 1.0 / hits as f32;
122        for a in acc.iter_mut() {
123            *a *= inv;
124        }
125        l2_normalize(&mut acc);
126        acc
127    }
128
129    fn dim(&self) -> usize {
130        self.dim
131    }
132
133    fn name(&self) -> &str {
134        &self.name
135    }
136}
137
138/// L2-normalize in place. Zero vectors are left as-is (cosine treats them
139/// as `0.0`).
140pub(crate) fn l2_normalize(vec: &mut [f32]) {
141    let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
142    if norm > 0.0 {
143        let inv = 1.0 / norm;
144        for x in vec.iter_mut() {
145            *x *= inv;
146        }
147    }
148}
149
150// --- minimal, dependency-free asset parser -------------------------------
151
152struct AssetDoc {
153    dim: usize,
154    vectors: HashMap<String, Vec<f32>>,
155}
156
157/// Parse the fixed `{ "dim": N, "vectors": { "tok": [floats] } }` shape.
158/// This avoids adding a JSON dependency to the default build for what is an
159/// optional asset; a future safetensors loader supersedes it entirely.
160fn parse_asset(raw: &str) -> Result<AssetDoc, String> {
161    // We lean on the harn-vm value layer? No — keep it standalone. Use a
162    // tiny tolerant scanner: find "dim": <int>, then "vectors": { ... }.
163    let dim = extract_int(raw, "\"dim\"")
164        .ok_or_else(|| "static embedding asset missing integer `dim`".to_string())?;
165    if dim == 0 {
166        return Err("static embedding `dim` must be > 0".to_string());
167    }
168    let vectors = extract_vectors(raw)?;
169    Ok(AssetDoc {
170        dim: dim as usize,
171        vectors,
172    })
173}
174
175fn extract_int(raw: &str, key: &str) -> Option<i64> {
176    let idx = raw.find(key)?;
177    let after = &raw[idx + key.len()..];
178    let colon = after.find(':')?;
179    let rest = after[colon + 1..].trim_start();
180    let end = rest
181        .find(|c: char| !c.is_ascii_digit() && c != '-')
182        .unwrap_or(rest.len());
183    rest[..end].parse::<i64>().ok()
184}
185
186fn extract_vectors(raw: &str) -> Result<HashMap<String, Vec<f32>>, String> {
187    let key = "\"vectors\"";
188    let idx = raw
189        .find(key)
190        .ok_or_else(|| "static embedding asset missing `vectors`".to_string())?;
191    let after = &raw[idx + key.len()..];
192    let open = after
193        .find('{')
194        .ok_or_else(|| "`vectors` is not an object".to_string())?;
195    let body = &after[open + 1..];
196    let mut map = HashMap::new();
197    let bytes = body.as_bytes();
198    let mut i = 0usize;
199    while i < bytes.len() {
200        // find next quote (start of a key) or closing brace of the object
201        match bytes[i] {
202            b'}' => break,
203            b'"' => {
204                let (k, next) = parse_string(body, i)?;
205                i = next;
206                // skip to colon
207                while i < bytes.len() && bytes[i] != b':' {
208                    i += 1;
209                }
210                i += 1;
211                // skip to array open
212                while i < bytes.len() && bytes[i] != b'[' {
213                    i += 1;
214                }
215                let (vec, next) = parse_float_array(body, i)?;
216                i = next;
217                map.insert(k, vec);
218            }
219            _ => i += 1,
220        }
221    }
222    Ok(map)
223}
224
225fn parse_string(s: &str, start: usize) -> Result<(String, usize), String> {
226    let bytes = s.as_bytes();
227    debug_assert_eq!(bytes[start], b'"');
228    let mut i = start + 1;
229    let mut out = String::new();
230    while i < bytes.len() {
231        match bytes[i] {
232            b'"' => return Ok((out, i + 1)),
233            b'\\' if i + 1 < bytes.len() => {
234                out.push(bytes[i + 1] as char);
235                i += 2;
236            }
237            c => {
238                out.push(c as char);
239                i += 1;
240            }
241        }
242    }
243    Err("unterminated string in static embedding asset".to_string())
244}
245
246fn parse_float_array(s: &str, start: usize) -> Result<(Vec<f32>, usize), String> {
247    let bytes = s.as_bytes();
248    if start >= bytes.len() || bytes[start] != b'[' {
249        return Err("expected float array in static embedding asset".to_string());
250    }
251    let mut i = start + 1;
252    let mut out = Vec::new();
253    let mut num = String::new();
254    let flush = |num: &mut String, out: &mut Vec<f32>| -> Result<(), String> {
255        let t = num.trim();
256        if !t.is_empty() {
257            out.push(
258                t.parse::<f32>()
259                    .map_err(|_| format!("bad float `{t}` in static embedding asset"))?,
260            );
261        }
262        num.clear();
263        Ok(())
264    };
265    while i < bytes.len() {
266        match bytes[i] {
267            b']' => {
268                flush(&mut num, &mut out)?;
269                return Ok((out, i + 1));
270            }
271            b',' => {
272                flush(&mut num, &mut out)?;
273                i += 1;
274            }
275            c if c.is_ascii_whitespace() => i += 1,
276            c => {
277                num.push(c as char);
278                i += 1;
279            }
280        }
281    }
282    Err("unterminated float array in static embedding asset".to_string())
283}
284
285/// Resolve the asset directory for a named embedding model, honoring an
286/// explicit override before falling back to a conventional location under
287/// the data dir. Returns `None` when nothing resolvable exists, which the
288/// caller treats as "use the lexical floor".
289///
290/// Resolution order (sandbox/settings-aware):
291/// 1. explicit `override_dir` (from a Harn setting / host call param),
292/// 2. `<data_dir>/embeddings/<model>` (conventional vendored location).
293///
294/// The function never touches the network and never reads outside the
295/// provided roots, so it is safe to call from inside a sandbox.
296pub fn resolve_asset_dir(
297    override_dir: Option<&Path>,
298    data_dir: Option<&Path>,
299    model: &str,
300) -> Option<PathBuf> {
301    if let Some(dir) = override_dir {
302        if dir.join("static-embeddings.json").is_file() {
303            return Some(dir.to_path_buf());
304        }
305    }
306    if let Some(base) = data_dir {
307        let candidate = base.join("embeddings").join(model);
308        if candidate.join("static-embeddings.json").is_file() {
309            return Some(candidate);
310        }
311    }
312    None
313}
314
315#[cfg(test)]
316mod tests {
317    use super::*;
318
319    #[test]
320    fn lexical_identical_text_is_self_similar() {
321        let e = LexicalEmbedder::default();
322        let v = e.embed("rate limiter middleware");
323        assert_eq!(v.len(), 256);
324        let sim = super::super::similarity::cosine(&v, &v);
325        assert!((sim - 1.0).abs() < 1e-5, "self-sim was {sim}");
326    }
327
328    #[test]
329    fn lexical_related_beats_unrelated() {
330        let e = LexicalEmbedder::default();
331        let query = e.embed("rate limiter for the API");
332        let related = e.embed("RateLimiter API throttle");
333        let unrelated = e.embed("parse markdown table renderer");
334        let s_rel = super::super::similarity::cosine(&query, &related);
335        let s_unrel = super::super::similarity::cosine(&query, &unrelated);
336        assert!(
337            s_rel > s_unrel,
338            "related {s_rel} should beat unrelated {s_unrel}"
339        );
340    }
341
342    #[test]
343    fn lexical_empty_is_zero_vector() {
344        let e = LexicalEmbedder::default();
345        let v = e.embed("");
346        assert!(v.iter().all(|&x| x == 0.0));
347    }
348
349    #[test]
350    fn lexical_is_l2_normalized() {
351        let e = LexicalEmbedder::default();
352        let v = e.embed("hello world embedding test");
353        let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
354        assert!((norm - 1.0).abs() < 1e-5, "norm was {norm}");
355    }
356
357    #[test]
358    fn lexical_is_deterministic_cross_run() {
359        // The whole cross-platform contract: same input -> same vector.
360        let e = LexicalEmbedder::default();
361        assert_eq!(e.embed("getUserById"), e.embed("getUserById"));
362    }
363
364    #[test]
365    fn static_embedder_pools_known_tokens() {
366        let json = r#"{ "dim": 2, "vectors": {
367            "rate": [1.0, 0.0],
368            "limit": [0.0, 1.0],
369            "throttle": [0.7071, 0.7071]
370        } }"#;
371        let e = StaticEmbedder::from_json(json).expect("parse");
372        assert_eq!(e.dim(), 2);
373        // "rate limit" pools (1,0)+(0,1) -> (0.5,0.5) -> normalized (1/sqrt2, 1/sqrt2)
374        let v = e.embed("rate limit");
375        let expected = std::f32::consts::FRAC_1_SQRT_2;
376        assert!((v[0] - expected).abs() < 1e-3, "{v:?}");
377        assert!((v[1] - expected).abs() < 1e-3, "{v:?}");
378        // "throttle" should be very close to "rate limit" semantically here.
379        let sim = super::super::similarity::cosine(&v, &e.embed("throttle"));
380        assert!(sim > 0.99, "throttle sim {sim}");
381    }
382
383    #[test]
384    fn static_embedder_falls_back_for_unknown_tokens() {
385        let json = r#"{ "dim": 2, "vectors": { "rate": [1.0, 0.0] } }"#;
386        let e = StaticEmbedder::from_json(json).expect("parse");
387        // Unknown tokens -> lexical fallback, non-zero, comparable.
388        let v = e.embed("zzz totally unknown words");
389        assert!(v.iter().any(|&x| x != 0.0));
390    }
391
392    #[test]
393    fn static_embedder_rejects_malformed_asset() {
394        assert!(StaticEmbedder::from_json("not json").is_err());
395        assert!(StaticEmbedder::from_json(r#"{ "dim": 2, "vectors": {} }"#).is_err());
396        // length mismatch
397        assert!(
398            StaticEmbedder::from_json(r#"{ "dim": 3, "vectors": { "x": [1.0, 2.0] } }"#).is_err()
399        );
400    }
401
402    #[test]
403    fn resolve_asset_dir_respects_override_and_absence() {
404        let tmp = std::env::temp_dir().join("embed-resolve-test-absent-xyz");
405        let _ = std::fs::remove_dir_all(&tmp);
406        assert_eq!(resolve_asset_dir(Some(&tmp), None, "potion"), None);
407        assert_eq!(resolve_asset_dir(None, Some(&tmp), "potion"), None);
408    }
409
410    #[test]
411    fn parse_handles_negative_and_scientific_floats() {
412        let json = r#"{ "dim": 3, "vectors": { "x": [-1.5, 0.0, 2.0] } }"#;
413        let e = StaticEmbedder::from_json(json).expect("parse");
414        assert_eq!(e.dim(), 3);
415    }
416}