use crate::embed::Embedder;
use sha2::{Digest, Sha256};
use std::collections::{HashMap, HashSet};
use std::io;
use std::path::Path;
pub fn cosine(a: &[f32], b: &[f32]) -> f32 {
let n = a.len().min(b.len());
let (mut dot, mut na, mut nb) = (0f32, 0f32, 0f32);
for i in 0..n {
dot += a[i] * b[i];
na += a[i] * a[i];
nb += b[i] * b[i];
}
if na == 0.0 || nb == 0.0 {
return 0.0;
}
dot / (na.sqrt() * nb.sqrt())
}
fn content_hash(model: &str, text: &str) -> String {
let mut h = Sha256::new();
h.update(model.as_bytes());
h.update([0u8]);
h.update(text.as_bytes());
format!("{:x}", h.finalize())
}
fn load_store(dir: &Path) -> io::Result<(usize, HashMap<String, usize>, Vec<f32>)> {
let idx_path = dir.join("index.json");
let blob_path = dir.join("vectors.f32");
let idx_exists = idx_path.is_file();
let blob_exists = blob_path.is_file();
if idx_exists != blob_exists {
eprintln!("kibble: ignoring partial vector store at {} (one of index.json/vectors.f32 is missing)", dir.display());
}
if !idx_exists || !blob_exists {
return Ok((0, HashMap::new(), Vec::new()));
}
let idx_raw = std::fs::read_to_string(&idx_path)?;
let v: serde_json::Value =
serde_json::from_str(&idx_raw).map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
let dim = v.get("dim").and_then(|d| d.as_u64()).unwrap_or(0) as usize;
let mut index = HashMap::new();
if let Some(entries) = v.get("entries").and_then(|e| e.as_object()) {
for (h, row) in entries {
if let Some(r) = row.as_u64() {
index.insert(h.clone(), r as usize);
}
}
}
let bytes = std::fs::read(&blob_path)?;
let mut blob = Vec::with_capacity(bytes.len() / 4);
for chunk in bytes.chunks_exact(4) {
blob.push(f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]));
}
Ok((dim, index, blob))
}
fn save_store(dir: &Path, dim: usize, index: &HashMap<String, usize>, blob: &[f32]) -> io::Result<()> {
std::fs::create_dir_all(dir)?;
let mut bytes = Vec::with_capacity(blob.len() * 4);
for f in blob {
bytes.extend_from_slice(&f.to_le_bytes());
}
std::fs::write(dir.join("vectors.f32"), &bytes)?;
let entries: serde_json::Map<String, serde_json::Value> =
index.iter().map(|(h, r)| (h.clone(), serde_json::json!(*r))).collect();
let idx = serde_json::json!({ "dim": dim, "entries": entries });
std::fs::write(dir.join("index.json"), serde_json::to_string(&idx).unwrap())?;
Ok(())
}
pub async fn get_or_embed<E: Embedder>(
embedder: &E,
store_dir: &Path,
model: &str,
texts: &[String],
batch_size: usize,
) -> io::Result<Vec<Vec<f32>>> {
let (mut dim, mut index, mut blob) = load_store(store_dir)?;
let hashes: Vec<String> = texts.iter().map(|t| content_hash(model, t)).collect();
let mut miss_texts: Vec<String> = Vec::new();
let mut miss_hashes: Vec<String> = Vec::new();
let mut pending: HashSet<String> = HashSet::new();
for (t, h) in texts.iter().zip(hashes.iter()) {
if !index.contains_key(h) && pending.insert(h.clone()) {
miss_texts.push(t.clone());
miss_hashes.push(h.clone());
}
}
let bs = batch_size.max(1);
let mut new_vecs: Vec<Vec<f32>> = Vec::with_capacity(miss_texts.len());
for chunk in miss_texts.chunks(bs) {
let mut vs = embedder.embed_batch(chunk).await?;
new_vecs.append(&mut vs);
}
for v in &new_vecs {
if dim == 0 {
dim = v.len();
} else if v.len() != dim {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("embedding dim mismatch: store={dim} got={}", v.len()),
));
}
}
for (h, v) in miss_hashes.iter().zip(new_vecs.iter()) {
if index.contains_key(h) {
continue;
}
let row = blob.len().checked_div(dim).unwrap_or(0);
blob.extend_from_slice(v);
index.insert(h.clone(), row);
}
if !miss_hashes.is_empty() {
save_store(store_dir, dim, &index, &blob)?;
}
let mut out = Vec::with_capacity(texts.len());
for h in &hashes {
let row = *index
.get(h)
.ok_or_else(|| io::Error::other("vector missing after embed"))?;
out.push(blob[row * dim..(row + 1) * dim].to_vec());
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::embed::{Embedder, StubEmbedder};
#[test]
fn cosine_basics() {
assert!((cosine(&[1.0, 0.0], &[1.0, 0.0]) - 1.0).abs() < 1e-6);
assert!(cosine(&[1.0, 0.0], &[0.0, 1.0]).abs() < 1e-6);
assert_eq!(cosine(&[0.0, 0.0], &[1.0, 1.0]), 0.0); }
struct CountingEmbedder {
inner: StubEmbedder,
calls: std::cell::RefCell<usize>,
}
impl Embedder for CountingEmbedder {
async fn embed_batch(&self, texts: &[String]) -> std::io::Result<Vec<Vec<f32>>> {
*self.calls.borrow_mut() += texts.len();
self.inner.embed_batch(texts).await
}
}
#[tokio::test]
async fn store_round_trips_and_caches_misses() {
let dir = std::env::temp_dir().join(format!("kibble_vec_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
let e = CountingEmbedder { inner: StubEmbedder::new(), calls: std::cell::RefCell::new(0) };
let texts = vec!["alpha".to_string(), "beta".to_string(), "alpha".to_string()];
let v1 = get_or_embed(&e, &dir, "m", &texts, 64).await.unwrap();
assert_eq!(v1.len(), 3);
assert_eq!(v1[0], v1[2]); assert_eq!(*e.calls.borrow(), 2);
let before = *e.calls.borrow();
let v2 = get_or_embed(&e, &dir, "m", &texts, 64).await.unwrap();
assert_eq!(*e.calls.borrow(), before); assert_eq!(v1, v2);
let _ = get_or_embed(&e, &dir, "other", &texts, 64).await.unwrap();
assert!(*e.calls.borrow() > before);
std::fs::remove_dir_all(&dir).ok();
}
}