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 (_, after) = raw.split_once(key)?;
177    let (_, rest) = after.split_once(':')?;
178    let rest = rest.trim_start();
179    let end = rest
180        .find(|c: char| !c.is_ascii_digit() && c != '-')
181        .unwrap_or(rest.len());
182    #[expect(
183        clippy::string_slice,
184        reason = "`end` comes from find on `rest`, a char boundary"
185    )]
186    let digits = &rest[..end];
187    digits.parse::<i64>().ok()
188}
189
190fn extract_vectors(raw: &str) -> Result<HashMap<String, Vec<f32>>, String> {
191    let key = "\"vectors\"";
192    let (_, after) = raw
193        .split_once(key)
194        .ok_or_else(|| "static embedding asset missing `vectors`".to_string())?;
195    let (_, body) = after
196        .split_once('{')
197        .ok_or_else(|| "`vectors` is not an object".to_string())?;
198    let mut map = HashMap::new();
199    let bytes = body.as_bytes();
200    let mut i = 0usize;
201    while i < bytes.len() {
202        // find next quote (start of a key) or closing brace of the object
203        match bytes[i] {
204            b'}' => break,
205            b'"' => {
206                let (k, next) = parse_string(body, i)?;
207                i = next;
208                // skip to colon
209                while i < bytes.len() && bytes[i] != b':' {
210                    i += 1;
211                }
212                i += 1;
213                // skip to array open
214                while i < bytes.len() && bytes[i] != b'[' {
215                    i += 1;
216                }
217                let (vec, next) = parse_float_array(body, i)?;
218                i = next;
219                map.insert(k, vec);
220            }
221            _ => i += 1,
222        }
223    }
224    Ok(map)
225}
226
227fn parse_string(s: &str, start: usize) -> Result<(String, usize), String> {
228    let bytes = s.as_bytes();
229    debug_assert_eq!(bytes[start], b'"');
230    let mut i = start + 1;
231    let mut out = String::new();
232    while i < bytes.len() {
233        match bytes[i] {
234            b'"' => return Ok((out, i + 1)),
235            b'\\' if i + 1 < bytes.len() => {
236                out.push(bytes[i + 1] as char);
237                i += 2;
238            }
239            c => {
240                out.push(c as char);
241                i += 1;
242            }
243        }
244    }
245    Err("unterminated string in static embedding asset".to_string())
246}
247
248fn parse_float_array(s: &str, start: usize) -> Result<(Vec<f32>, usize), String> {
249    let bytes = s.as_bytes();
250    if start >= bytes.len() || bytes[start] != b'[' {
251        return Err("expected float array in static embedding asset".to_string());
252    }
253    let mut i = start + 1;
254    let mut out = Vec::new();
255    let mut num = String::new();
256    let flush = |num: &mut String, out: &mut Vec<f32>| -> Result<(), String> {
257        let t = num.trim();
258        if !t.is_empty() {
259            out.push(
260                t.parse::<f32>()
261                    .map_err(|_| format!("bad float `{t}` in static embedding asset"))?,
262            );
263        }
264        num.clear();
265        Ok(())
266    };
267    while i < bytes.len() {
268        match bytes[i] {
269            b']' => {
270                flush(&mut num, &mut out)?;
271                return Ok((out, i + 1));
272            }
273            b',' => {
274                flush(&mut num, &mut out)?;
275                i += 1;
276            }
277            c if c.is_ascii_whitespace() => i += 1,
278            c => {
279                num.push(c as char);
280                i += 1;
281            }
282        }
283    }
284    Err("unterminated float array in static embedding asset".to_string())
285}
286
287/// Resolve the asset directory for a named embedding model, honoring an
288/// explicit override before falling back to a conventional location under
289/// the data dir. Returns `None` when nothing resolvable exists, which the
290/// caller treats as "use the lexical floor".
291///
292/// Resolution order (sandbox/settings-aware):
293/// 1. explicit `override_dir` (from a Harn setting / host call param),
294/// 2. `<data_dir>/embeddings/<model>` (conventional vendored location).
295///
296/// The function never touches the network and never reads outside the
297/// provided roots, so it is safe to call from inside a sandbox.
298pub fn resolve_asset_dir(
299    override_dir: Option<&Path>,
300    data_dir: Option<&Path>,
301    model: &str,
302) -> Option<PathBuf> {
303    if let Some(dir) = override_dir {
304        if dir.join("static-embeddings.json").is_file() {
305            return Some(dir.to_path_buf());
306        }
307    }
308    if let Some(base) = data_dir {
309        let candidate = base.join("embeddings").join(model);
310        if candidate.join("static-embeddings.json").is_file() {
311            return Some(candidate);
312        }
313    }
314    None
315}
316
317#[cfg(test)]
318mod tests {
319    use super::*;
320
321    #[test]
322    fn lexical_identical_text_is_self_similar() {
323        let e = LexicalEmbedder::default();
324        let v = e.embed("rate limiter middleware");
325        assert_eq!(v.len(), 256);
326        let sim = super::super::similarity::cosine(&v, &v);
327        assert!((sim - 1.0).abs() < 1e-5, "self-sim was {sim}");
328    }
329
330    #[test]
331    fn lexical_related_beats_unrelated() {
332        let e = LexicalEmbedder::default();
333        let query = e.embed("rate limiter for the API");
334        let related = e.embed("RateLimiter API throttle");
335        let unrelated = e.embed("parse markdown table renderer");
336        let s_rel = super::super::similarity::cosine(&query, &related);
337        let s_unrel = super::super::similarity::cosine(&query, &unrelated);
338        assert!(
339            s_rel > s_unrel,
340            "related {s_rel} should beat unrelated {s_unrel}"
341        );
342    }
343
344    #[test]
345    fn lexical_empty_is_zero_vector() {
346        let e = LexicalEmbedder::default();
347        let v = e.embed("");
348        assert!(v.iter().all(|&x| x == 0.0));
349    }
350
351    #[test]
352    fn lexical_is_l2_normalized() {
353        let e = LexicalEmbedder::default();
354        let v = e.embed("hello world embedding test");
355        let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
356        assert!((norm - 1.0).abs() < 1e-5, "norm was {norm}");
357    }
358
359    #[test]
360    fn lexical_is_deterministic_cross_run() {
361        // The whole cross-platform contract: same input -> same vector.
362        let e = LexicalEmbedder::default();
363        assert_eq!(e.embed("getUserById"), e.embed("getUserById"));
364    }
365
366    #[test]
367    fn static_embedder_pools_known_tokens() {
368        let json = r#"{ "dim": 2, "vectors": {
369            "rate": [1.0, 0.0],
370            "limit": [0.0, 1.0],
371            "throttle": [0.7071, 0.7071]
372        } }"#;
373        let e = StaticEmbedder::from_json(json).expect("parse");
374        assert_eq!(e.dim(), 2);
375        // "rate limit" pools (1,0)+(0,1) -> (0.5,0.5) -> normalized (1/sqrt2, 1/sqrt2)
376        let v = e.embed("rate limit");
377        let expected = std::f32::consts::FRAC_1_SQRT_2;
378        assert!((v[0] - expected).abs() < 1e-3, "{v:?}");
379        assert!((v[1] - expected).abs() < 1e-3, "{v:?}");
380        // "throttle" should be very close to "rate limit" semantically here.
381        let sim = super::super::similarity::cosine(&v, &e.embed("throttle"));
382        assert!(sim > 0.99, "throttle sim {sim}");
383    }
384
385    #[test]
386    fn static_embedder_falls_back_for_unknown_tokens() {
387        let json = r#"{ "dim": 2, "vectors": { "rate": [1.0, 0.0] } }"#;
388        let e = StaticEmbedder::from_json(json).expect("parse");
389        // Unknown tokens -> lexical fallback, non-zero, comparable.
390        let v = e.embed("zzz totally unknown words");
391        assert!(v.iter().any(|&x| x != 0.0));
392    }
393
394    #[test]
395    fn static_embedder_rejects_malformed_asset() {
396        assert!(StaticEmbedder::from_json("not json").is_err());
397        assert!(StaticEmbedder::from_json(r#"{ "dim": 2, "vectors": {} }"#).is_err());
398        // length mismatch
399        assert!(
400            StaticEmbedder::from_json(r#"{ "dim": 3, "vectors": { "x": [1.0, 2.0] } }"#).is_err()
401        );
402    }
403
404    #[test]
405    fn resolve_asset_dir_respects_override_and_absence() {
406        let tmp = std::env::temp_dir().join("embed-resolve-test-absent-xyz");
407        let _ = std::fs::remove_dir_all(&tmp);
408        assert_eq!(resolve_asset_dir(Some(&tmp), None, "potion"), None);
409        assert_eq!(resolve_asset_dir(None, Some(&tmp), "potion"), None);
410    }
411
412    #[test]
413    fn parse_handles_negative_and_scientific_floats() {
414        let json = r#"{ "dim": 3, "vectors": { "x": [-1.5, 0.0, 2.0] } }"#;
415        let e = StaticEmbedder::from_json(json).expect("parse");
416        assert_eq!(e.dim(), 3);
417    }
418}