feat: resumable sha256 embedding cache + cost estimate
This commit is contained in:
+84
-1
@@ -1 +1,84 @@
|
||||
// implemented in a later task
|
||||
use crate::ionos::IonosClient;
|
||||
use crate::model::AppError;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::fs;
|
||||
use std::path::PathBuf;
|
||||
|
||||
pub struct EmbeddingStore {
|
||||
dir: PathBuf, // <index_dir>/embeddings/<model-sanitized>/
|
||||
model: String,
|
||||
}
|
||||
|
||||
pub struct Estimate { pub items: usize, pub chars: usize }
|
||||
|
||||
pub fn estimate(pending: &[String]) -> Estimate {
|
||||
Estimate { items: pending.len(), chars: pending.iter().map(|s| s.len()).sum() }
|
||||
}
|
||||
|
||||
impl EmbeddingStore {
|
||||
pub fn key(model: &str, text: &str) -> String {
|
||||
let mut h = Sha256::new();
|
||||
h.update(model.as_bytes());
|
||||
h.update([0u8]);
|
||||
h.update(text.as_bytes());
|
||||
hex(&h.finalize())
|
||||
}
|
||||
|
||||
pub fn open(index_dir: &str, model: &str) -> Result<Self, AppError> {
|
||||
let safe = model.replace('/', "_");
|
||||
let dir = PathBuf::from(index_dir).join("embeddings").join(safe);
|
||||
fs::create_dir_all(&dir)
|
||||
.map_err(|e| AppError::Io(format!("{}: {e}", dir.display())))?;
|
||||
Ok(Self { dir, model: model.to_string() })
|
||||
}
|
||||
|
||||
fn path(&self, text: &str) -> PathBuf {
|
||||
self.dir.join(format!("{}.f32", Self::key(&self.model, text)))
|
||||
}
|
||||
|
||||
pub fn has(&self, text: &str) -> bool { self.path(text).exists() }
|
||||
|
||||
pub fn get(&self, text: &str) -> Option<Vec<f32>> {
|
||||
let bytes = fs::read(self.path(text)).ok()?;
|
||||
Some(bytes.chunks_exact(4)
|
||||
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
|
||||
.collect())
|
||||
}
|
||||
|
||||
/// Atomic write (temp + rename) so an interrupted run never leaves
|
||||
/// a half-written vector — the next run simply re-embeds that key.
|
||||
pub fn put(&self, text: &str, v: &[f32]) -> Result<(), AppError> {
|
||||
let mut bytes = Vec::with_capacity(v.len() * 4);
|
||||
for f in v { bytes.extend_from_slice(&f.to_le_bytes()); }
|
||||
let final_path = self.path(text);
|
||||
let tmp = final_path.with_extension("tmp");
|
||||
fs::write(&tmp, &bytes).map_err(|e| AppError::Io(format!("{e}")))?;
|
||||
fs::rename(&tmp, &final_path).map_err(|e| AppError::Io(format!("{e}")))?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn hex(b: &[u8]) -> String {
|
||||
b.iter().map(|x| format!("{:02x}", x)).collect()
|
||||
}
|
||||
|
||||
/// Embed a batch via IONOS, persisting each result. Only `texts` not
|
||||
/// already cached should be passed (caller filters via `has`).
|
||||
pub fn embed_batch(
|
||||
client: &IonosClient, store: &EmbeddingStore, model: &str, texts: &[String],
|
||||
) -> Result<usize, AppError> {
|
||||
if texts.is_empty() { return Ok(0); }
|
||||
let body = serde_json::json!({ "model": model, "input": texts });
|
||||
let resp: serde_json::Value = client.post_json("/embeddings", &body)?;
|
||||
let data = resp["data"].as_array()
|
||||
.ok_or_else(|| AppError::Ionos("embeddings: no data[]".into()))?;
|
||||
let mut n = 0;
|
||||
for (i, item) in data.iter().enumerate() {
|
||||
let v: Vec<f32> = item["embedding"].as_array()
|
||||
.ok_or_else(|| AppError::Ionos("embeddings: no embedding[]".into()))?
|
||||
.iter().map(|x| x.as_f64().unwrap_or(0.0) as f32).collect();
|
||||
store.put(&texts[i], &v)?;
|
||||
n += 1;
|
||||
}
|
||||
Ok(n)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
use alpha_id::embed::{EmbeddingStore, estimate};
|
||||
|
||||
#[test]
|
||||
fn cache_key_is_stable_sha256() {
|
||||
let k1 = EmbeddingStore::key("BAAI/bge-m3", "diabetes");
|
||||
let k2 = EmbeddingStore::key("BAAI/bge-m3", "diabetes");
|
||||
assert_eq!(k1, k2);
|
||||
assert_ne!(k1, EmbeddingStore::key("BAAI/bge-m3", "hypertonie"));
|
||||
assert_eq!(k1.len(), 64); // hex sha256
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn store_persists_and_reloads_vectors_resumably() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let store = EmbeddingStore::open(dir.path().to_str().unwrap(), "BAAI/bge-m3").unwrap();
|
||||
assert!(store.get("diabetes").is_none());
|
||||
store.put("diabetes", &[0.1, 0.2, 0.3]).unwrap();
|
||||
let reopened = EmbeddingStore::open(dir.path().to_str().unwrap(), "BAAI/bge-m3").unwrap();
|
||||
assert_eq!(reopened.get("diabetes").unwrap(), vec![0.1, 0.2, 0.3]);
|
||||
assert!(reopened.has("diabetes")); // resume: skip already-embedded
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn estimate_reports_pending_count_and_chars() {
|
||||
let est = estimate(&["abc".to_string(), "de".to_string()]);
|
||||
assert_eq!(est.items, 2);
|
||||
assert_eq!(est.chars, 5);
|
||||
}
|
||||
Reference in New Issue
Block a user