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 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 match bytes[i] {
202 b'}' => break,
203 b'"' => {
204 let (k, next) = parse_string(body, i)?;
205 i = next;
206 while i < bytes.len() && bytes[i] != b':' {
208 i += 1;
209 }
210 i += 1;
211 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
285pub 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 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 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 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 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 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}