1use std::collections::HashMap;
28use std::path::{Path, PathBuf};
29
30pub use harn_session_store::{Embedder, LexicalEmbedder};
31
32use super::tokenize;
33
34pub struct StaticEmbedder {
53 dim: usize,
54 vectors: HashMap<String, Vec<f32>>,
55 name: String,
56 fallback: LexicalEmbedder,
60}
61
62impl StaticEmbedder {
63 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 pub fn from_json(raw: &str) -> Result<Self, String> {
79 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 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
138pub(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
150struct AssetDoc {
153 dim: usize,
154 vectors: HashMap<String, Vec<f32>>,
155}
156
157fn parse_asset(raw: &str) -> Result<AssetDoc, String> {
161 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 match bytes[i] {
204 b'}' => break,
205 b'"' => {
206 let (k, next) = parse_string(body, i)?;
207 i = next;
208 while i < bytes.len() && bytes[i] != b':' {
210 i += 1;
211 }
212 i += 1;
213 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
287pub 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 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 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 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 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 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}