Skip to main content

rig_core/embeddings/
builder.rs

1//! Batched embedding generation for documents implementing [`Embed`].
2//!
3//! ```no_run
4//! use rig_core::DynModel;
5//! use rig_core::embeddings::EmbeddingsBuilder;
6//! use rig_core::operation::Embedding;
7//!
8//! # async fn example(model: DynModel<Embedding>) -> Result<(), Box<dyn std::error::Error>> {
9//! let documents = EmbeddingsBuilder::new(model)
10//!     .documents(["first document", "second document"])?
11//!     .build().await?;
12//! # let _ = documents;
13//! # Ok(())
14//! # }
15//! ```
16
17use std::{cmp::max, ops::Range};
18
19use futures::{StreamExt, stream};
20
21use crate::driver::DynModel;
22use crate::error::ProviderError;
23use crate::operation::Embedding as EmbeddingOp;
24use crate::{
25    completion::Usage,
26    embeddings::{Embed, EmbedError, Embedding, EmbeddingResponse, embed::TextEmbedder},
27};
28
29/// Accumulates documents and embeds their extracted texts in provider-sized batches.
30/// Text extraction occurs when documents are added; requests start when built.
31#[must_use = "an embeddings builder does nothing until built"]
32pub struct EmbeddingsBuilder<T> {
33    model: DynModel<EmbeddingOp>,
34    documents: Vec<(T, Vec<String>)>,
35}
36
37impl<T: Embed> EmbeddingsBuilder<T> {
38    /// Create a new embedding builder with the given embedding model
39    pub fn new(model: impl Into<DynModel<EmbeddingOp>>) -> Self {
40        Self {
41            model: model.into(),
42            documents: vec![],
43        }
44    }
45
46    /// Add a document to be embedded to the builder. `document` must implement the [Embed] trait.
47    pub fn document(mut self, document: T) -> Result<Self, EmbedError> {
48        let mut embedder = TextEmbedder::default();
49        document.embed(&mut embedder)?;
50
51        self.documents.push((document, embedder.texts));
52
53        Ok(self)
54    }
55
56    /// Add multiple documents to be embedded to the builder. `documents` must be iterable
57    /// with items that implement the [Embed] trait.
58    pub fn documents(self, documents: impl IntoIterator<Item = T>) -> Result<Self, EmbedError> {
59        let builder = documents
60            .into_iter()
61            .try_fold(self, EmbeddingsBuilder::document)?;
62
63        Ok(builder)
64    }
65}
66
67impl<T> EmbeddingsBuilder<T>
68where
69    T: Embed + crate::wasm_compat::WasmCompatSend,
70{
71    /// Generate embeddings for all documents in the builder.
72    ///
73    /// Returns `(document, embeddings)` pairs. A document may produce one or many
74    /// embeddings depending on how its [`Embed`] implementation uses [`TextEmbedder`].
75    ///
76    /// Preserves document insertion order and each document's text order,
77    /// regardless of batch completion order. Providers must return embeddings
78    /// in input order within each batch; this is not checked.
79    ///
80    /// Propagates provider and transport errors. Returns an error identifying
81    /// the document if it produces no text or a batch returns too few embeddings.
82    /// Empty embedded collections produce no text. Surplus embeddings are ignored.
83    pub async fn build(self) -> Result<Vec<(T, Vec<Embedding>)>, ProviderError> {
84        let (result, _usage) = self.build_with_usage().await?;
85        Ok(result)
86    }
87
88    /// Generate embeddings for all documents in the builder and return accumulated token usage.
89    ///
90    /// Returns `(document, embeddings)` pairs and the total token usage across all
91    /// batches. A document may produce one or many embeddings depending on how its
92    /// [`Embed`] implementation uses [`TextEmbedder`].
93    ///
94    /// Ordering is guaranteed at both levels, and the same two errors originate
95    /// here; both are described on [`Self::build`].
96    pub(crate) async fn build_with_usage(
97        self,
98    ) -> Result<(Vec<(T, Vec<Embedding>)>, Usage), ProviderError> {
99        use stream::TryStreamExt;
100
101        // Per-text slots preserve order even when a document spans batches
102        // that finish out of order.
103        let mut docs: Vec<T> = Vec::with_capacity(self.documents.len());
104        let mut spans: Vec<Range<usize>> = Vec::with_capacity(self.documents.len());
105        let mut texts: Vec<String> = Vec::new();
106
107        for (doc, doc_texts) in self.documents {
108            let start = texts.len();
109            texts.extend(doc_texts);
110            spans.push(start..texts.len());
111            docs.push(doc);
112        }
113
114        let total_texts = texts.len();
115        let max_documents = max(1, self.model.capabilities().max_documents);
116
117        let (slots, usage) = stream::iter(texts.into_iter().enumerate())
118            .chunks(max_documents)
119            .map(|chunk| async {
120                let (slots, batch): (Vec<usize>, Vec<String>) = chunk.into_iter().unzip();
121
122                let response: EmbeddingResponse = self.model.call(batch).await?;
123                Ok::<_, ProviderError>((
124                    slots
125                        .into_iter()
126                        .zip(response.embeddings)
127                        .collect::<Vec<_>>(),
128                    response.usage,
129                ))
130            })
131            .buffer_unordered(max(1, 1024 / max_documents))
132            .try_fold(
133                (
134                    (0..total_texts)
135                        .map(|_| None)
136                        .collect::<Vec<Option<Embedding>>>(),
137                    Usage::default(),
138                ),
139                |(mut slots, mut usage_acc), (chunk_embeddings, chunk_usage)| async move {
140                    for (slot, embedding) in chunk_embeddings {
141                        // Enumerated slots remain in range; zip discards any
142                        // surplus provider embeddings.
143                        if let Some(place) = slots.get_mut(slot) {
144                            *place = Some(embedding);
145                        }
146                    }
147                    usage_acc += chunk_usage;
148                    Ok((slots, usage_acc))
149                },
150            )
151            .await?;
152
153        let mut slots = slots.into_iter();
154        let mut result = Vec::with_capacity(docs.len());
155
156        for (index, (doc, span)) in docs.into_iter().zip(spans).enumerate() {
157            if span.is_empty() {
158                return Err(crate::error::ProviderError::Response(format!(
159                    "document {index} produced no text to embed, so it has no \
160                     embeddings to return; an empty collection in an `#[embed]` \
161                     field embeds nothing"
162                )));
163            }
164
165            // Missing slots identify short provider responses without silently
166            // dropping a document's texts.
167            let embeddings = slots
168                .by_ref()
169                .take(span.len())
170                .collect::<Option<Vec<Embedding>>>()
171                .ok_or_else(|| {
172                    crate::error::ProviderError::Response(format!(
173                        "provider returned fewer embeddings than texts sent: \
174                         document {index} is missing at least one of its {} texts \
175                         (slots {}..{} of {total_texts})",
176                        span.len(),
177                        span.start,
178                        span.end
179                    ))
180                })?;
181
182            result.push((doc, embeddings));
183        }
184
185        Ok((result, usage))
186    }
187}
188
189#[cfg(test)]
190mod tests;