use super::{
Embed, EmbedError, Embedding, EmbeddingError, EmbeddingModel,
embed::TextEmbedder,
};
pub struct EmbeddingsBuilder<M, T> {
model: M,
documents: Vec<(T, Vec<String>)>,
}
impl<M: EmbeddingModel, T: Embed> EmbeddingsBuilder<M, T> {
pub fn new(model: M) -> Self {
Self { model, documents: Vec::new() }
}
pub fn document(mut self, doc: T) -> Result<Self, EmbedError> {
let mut embedder = TextEmbedder::default();
doc.embed(&mut embedder)?;
self.documents.push((doc, embedder.texts));
Ok(self)
}
pub fn documents(self, docs: impl IntoIterator<Item = T>) -> Result<Self, EmbedError> {
docs.into_iter().try_fold(self, |b, doc| b.document(doc))
}
pub async fn build(self) -> Result<Vec<(T, Vec<Embedding>)>, EmbeddingError> {
let mut flat: Vec<(usize, String)> = Vec::new();
for (i, (_, texts)) in self.documents.iter().enumerate() {
for text in texts {
flat.push((i, text.clone()));
}
}
let mut embeddings_by_doc: Vec<Vec<Embedding>> =
(0..self.documents.len()).map(|_| Vec::new()).collect();
for chunk in flat.chunks(M::MAX_DOCUMENTS) {
let (ids, texts): (Vec<usize>, Vec<String>) = chunk
.iter()
.cloned()
.unzip();
let batch = self
.model
.embed_texts(texts)
.await
.map_err(|e| EmbeddingError::Response(e.to_string()))?;
if batch.len() != ids.len() {
return Err(EmbeddingError::Response(format!(
"model returned {} embeddings for {} inputs",
batch.len(),
ids.len(),
)));
}
for (doc_idx, embedding) in ids.into_iter().zip(batch) {
embeddings_by_doc[doc_idx].push(embedding);
}
}
let result = self
.documents
.into_iter()
.zip(embeddings_by_doc)
.map(|((doc, _), embeddings)| (doc, embeddings))
.collect();
Ok(result)
}
}