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;