use crate::Result;
pub const DEFAULT_EMBEDDING_DIM: usize = 384;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(transparent))]
pub struct Embedding(pub Vec<f32>);
impl Embedding {
pub fn dim(&self) -> usize {
self.0.len()
}
pub fn as_slice(&self) -> &[f32] {
&self.0
}
}
pub trait Embedder: Send + Sync {
fn dim(&self) -> usize;
fn embed(&self, text: &str) -> Result<Embedding>;
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
texts.iter().map(|t| self.embed(t)).collect()
}
fn model_id(&self) -> &str {
"unknown"
}
}
#[cfg(test)]
mod tests {
use super::*;
struct ConstEmbedder;
impl Embedder for ConstEmbedder {
fn dim(&self) -> usize {
2
}
fn embed(&self, text: &str) -> Result<Embedding> {
Ok(Embedding(vec![text.len() as f32, 0.0]))
}
}
#[test]
fn default_dim_matches_mempalace_for_migration() {
assert_eq!(DEFAULT_EMBEDDING_DIM, 384);
}
#[test]
fn embed_batch_defaults_to_per_item_loop() {
let e = ConstEmbedder;
let got = e.embed_batch(&["a", "bb", "ccc"]).expect("must embed");
assert_eq!(got.len(), 3);
assert_eq!(got[0].dim(), 2);
assert_eq!(got[2].as_slice(), &[3.0, 0.0]);
}
}