Skip to main content

ferrin_core/
embed.rs

1//! Embeddings: [`embed`] for one value, [`embed_many`] for several (split
2//! into provider calls by the model's limits) and [`cosine_similarity`].
3//!
4//! Design: `docs/01-architecture/11-other-modalities.md` ยง1.
5
6use std::future::IntoFuture;
7use std::sync::Arc;
8use std::time::Instant;
9
10use ferrin_spec::BoxFuture;
11use ferrin_spec::DynEmbeddingModel;
12use ferrin_spec::EmbeddingModelRef;
13use ferrin_spec::Headers;
14use ferrin_spec::ProviderMetadata;
15use ferrin_spec::ProviderOptions;
16use ferrin_spec::ResponseMetadata;
17use ferrin_spec::Warning;
18use ferrin_spec::embedding_model::EmbedOptions;
19use ferrin_spec::embedding_model::EmbedResult as ModelEmbedResult;
20pub use ferrin_spec::embedding_model::Embedding;
21use ferrin_spec::error::InvalidResponseDataError;
22use ferrin_spec::error::ProviderError;
23use serde_json::json;
24use tokio::task::JoinSet;
25use tokio_util::sync::CancellationToken;
26use tracing::Instrument;
27
28use crate::error::Error;
29use crate::ids::default_id_generator;
30use crate::modality::ModalityOptions;
31use crate::modality::accumulate_provider_metadata;
32use crate::modality::impl_modality_builder;
33use crate::registry::ProviderRegistry;
34use crate::registry::default::resolve_model;
35use crate::retry::RetryPolicy;
36use crate::retry::retry;
37use crate::telemetry::EmbedEndEvent;
38use crate::telemetry::EmbedStartEvent;
39use crate::telemetry::ErrorEvent;
40use crate::telemetry::ErrorPhase;
41use crate::telemetry::ModelIdentity;
42use crate::telemetry::dispatcher::TelemetryDispatcher;
43use crate::telemetry::spans;
44
45/// Token usage of embedding calls; `None` when the provider reported none.
46#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
47pub struct EmbeddingUsage {
48    /// Input tokens consumed.
49    pub tokens: Option<u64>,
50}
51
52/// Result of [`embed`].
53#[derive(Debug, Clone, PartialEq)]
54pub struct EmbedResult {
55    /// The embedded value.
56    pub value: String,
57    /// The embedding.
58    pub embedding: Embedding,
59    /// Token usage.
60    pub usage: EmbeddingUsage,
61    /// Adapter warnings.
62    pub warnings: Vec<Warning>,
63    /// Response metadata.
64    pub response: ResponseMetadata,
65    /// Provider-specific metadata.
66    pub provider_metadata: Option<ProviderMetadata>,
67}
68
69/// Result of [`embed_many`].
70#[derive(Debug, Clone, PartialEq)]
71pub struct EmbedManyResult {
72    /// The embedded values, in input order.
73    pub values: Vec<String>,
74    /// One embedding per value, in input order.
75    pub embeddings: Vec<Embedding>,
76    /// Token usage summed over all calls.
77    pub usage: EmbeddingUsage,
78    /// Adapter warnings of all calls.
79    pub warnings: Vec<Warning>,
80    /// Response metadata of every call.
81    pub responses: Vec<ResponseMetadata>,
82    /// Provider-specific metadata merged over all calls.
83    pub provider_metadata: Option<ProviderMetadata>,
84}
85
86/// Embeds one value.
87#[must_use]
88pub fn embed(model: impl Into<EmbeddingModelRef>, value: impl Into<String>) -> Embed {
89    Embed {
90        model: model.into(),
91        value: value.into(),
92        base: ModalityOptions::default(),
93    }
94}
95
96/// Builder returned by [`embed`]; `.await` runs the call.
97#[derive(Debug)]
98pub struct Embed {
99    model: EmbeddingModelRef,
100    value: String,
101    base: ModalityOptions,
102}
103
104impl_modality_builder!(Embed);
105
106impl IntoFuture for Embed {
107    type Output = Result<EmbedResult, Error>;
108    type IntoFuture = BoxFuture<'static, Self::Output>;
109
110    fn into_future(self) -> Self::IntoFuture {
111        Box::pin(async move {
112            let many = run(self.model, vec![self.value], self.base, Some(1)).await?;
113            let EmbedManyResult {
114                values,
115                embeddings,
116                usage,
117                warnings,
118                responses,
119                provider_metadata,
120            } = many;
121            let (Some(value), Some(embedding), Some(response)) = (
122                values.into_iter().next(),
123                embeddings.into_iter().next(),
124                responses.into_iter().next(),
125            ) else {
126                return Err(invalid_count(1, 0));
127            };
128            Ok(EmbedResult {
129                value,
130                embedding,
131                usage,
132                warnings,
133                response,
134                provider_metadata,
135            })
136        })
137    }
138}
139
140/// Embeds several values, splitting them into provider calls by the
141/// model's limits and running the calls concurrently when the model
142/// allows it.
143#[must_use]
144pub fn embed_many(
145    model: impl Into<EmbeddingModelRef>,
146    values: impl IntoIterator<Item = impl Into<String>>,
147) -> EmbedMany {
148    EmbedMany {
149        model: model.into(),
150        values: values.into_iter().map(Into::into).collect(),
151        max_parallel_calls: None,
152        base: ModalityOptions::default(),
153    }
154}
155
156/// Builder returned by [`embed_many`]; `.await` runs the calls.
157#[derive(Debug)]
158pub struct EmbedMany {
159    model: EmbeddingModelRef,
160    values: Vec<String>,
161    max_parallel_calls: Option<usize>,
162    base: ModalityOptions,
163}
164
165impl EmbedMany {
166    /// Limits the number of concurrent provider calls (default: unlimited;
167    /// ignored when the model does not support parallel calls).
168    #[must_use]
169    pub fn max_parallel_calls(mut self, max_parallel_calls: usize) -> Self {
170        self.max_parallel_calls = Some(max_parallel_calls.max(1));
171        self
172    }
173}
174
175impl_modality_builder!(EmbedMany);
176
177impl IntoFuture for EmbedMany {
178    type Output = Result<EmbedManyResult, Error>;
179    type IntoFuture = BoxFuture<'static, Self::Output>;
180
181    fn into_future(self) -> Self::IntoFuture {
182        Box::pin(run(
183            self.model,
184            self.values,
185            self.base,
186            self.max_parallel_calls,
187        ))
188    }
189}
190
191/// Cosine similarity of two vectors; `0.0` when either has zero length.
192///
193/// # Errors
194///
195/// Returns [`Error::InvalidArgument`] when the vectors have different
196/// dimensions.
197pub fn cosine_similarity(a: &[f32], b: &[f32]) -> Result<f32, Error> {
198    if a.len() != b.len() {
199        return Err(Error::invalid_argument(
200            "vectors",
201            format!(
202                "vectors must have the same length (got {} and {})",
203                a.len(),
204                b.len()
205            ),
206        ));
207    }
208    let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
209    let norm_a = a.iter().map(|x| x * x).sum::<f32>().sqrt();
210    let norm_b = b.iter().map(|y| y * y).sum::<f32>().sqrt();
211    if norm_a == 0.0 || norm_b == 0.0 {
212        return Ok(0.0);
213    }
214    Ok(dot / (norm_a * norm_b))
215}
216
217/// Splits `values` into chunks that respect both limits: a chunk is closed
218/// when it already holds `max_embeddings` values or when adding the next
219/// value would exceed `max_bytes`; a single oversized value still forms
220/// its own chunk.
221pub(crate) fn split_by_limits(
222    values: &[String],
223    max_embeddings: usize,
224    max_bytes: usize,
225) -> Vec<Vec<String>> {
226    let mut chunks: Vec<Vec<String>> = Vec::new();
227    let mut current: Vec<String> = Vec::new();
228    let mut current_bytes = 0usize;
229    for value in values {
230        let bytes = value.len();
231        if !current.is_empty()
232            && (current.len() >= max_embeddings || current_bytes.saturating_add(bytes) > max_bytes)
233        {
234            chunks.push(std::mem::take(&mut current));
235            current_bytes = 0;
236        }
237        current.push(value.clone());
238        current_bytes = current_bytes.saturating_add(bytes);
239    }
240    if !current.is_empty() {
241        chunks.push(current);
242    }
243    chunks
244}
245
246fn invalid_count(expected: usize, received: usize) -> Error {
247    Error::from(ProviderError::InvalidResponseData(Box::new(
248        InvalidResponseDataError::new(
249            format!("expected {expected} embeddings, received {received}"),
250            json!({ "expected": expected, "received": received }),
251        ),
252    )))
253}
254
255/// Everything one chunk call needs (owned, so it can run on a task).
256struct ChunkCall {
257    model: Arc<dyn DynEmbeddingModel>,
258    identity: ModelIdentity,
259    values: Vec<String>,
260    headers: Headers,
261    provider_options: ProviderOptions,
262    retry_policy: RetryPolicy,
263    cancellation: CancellationToken,
264    telemetry: TelemetryDispatcher,
265    call_id: String,
266}
267
268impl ChunkCall {
269    async fn run(self) -> Result<ModelEmbedResult, Error> {
270        let Self {
271            model,
272            identity,
273            values,
274            headers,
275            provider_options,
276            retry_policy,
277            cancellation,
278            telemetry,
279            call_id,
280        } = self;
281        let outcome = retry(&retry_policy, &cancellation, |_| {
282            let values = values.clone();
283            let model = &model;
284            let identity = &identity;
285            let headers = &headers;
286            let provider_options = &provider_options;
287            let cancellation = &cancellation;
288            let telemetry = &telemetry;
289            let call_id = &call_id;
290            async move {
291                let started = Instant::now();
292                telemetry.on_embed_start(&EmbedStartEvent {
293                    call_id: call_id.clone(),
294                    model: identity.clone(),
295                    value_count: values.len(),
296                    values: telemetry.record_inputs().then(|| values.clone()),
297                });
298                let result = model
299                    .do_embed(EmbedOptions {
300                        values,
301                        headers: headers.clone(),
302                        provider_options: provider_options.clone(),
303                        cancellation: cancellation.child_token(),
304                    })
305                    .await
306                    .map_err(Error::from)?;
307                telemetry.on_embed_end(&EmbedEndEvent {
308                    call_id: call_id.clone(),
309                    embedding_count: result.embeddings.len(),
310                    tokens: result.usage.map(|usage| usage.tokens),
311                    duration: started.elapsed(),
312                });
313                Ok(result)
314            }
315        })
316        .await;
317        let result = match outcome {
318            Ok(result) => result,
319            Err(error) => {
320                telemetry.on_error(&ErrorEvent {
321                    call_id: &call_id,
322                    error: &error,
323                    phase: ErrorPhase::ModelCall,
324                });
325                return Err(error);
326            }
327        };
328        if result.embeddings.len() != values.len() {
329            return Err(invalid_count(values.len(), result.embeddings.len()));
330        }
331        spans::log_warnings(&result.warnings, &identity);
332        Ok(result)
333    }
334}
335
336async fn run(
337    model: EmbeddingModelRef,
338    values: Vec<String>,
339    base: ModalityOptions,
340    max_parallel_calls: Option<usize>,
341) -> Result<EmbedManyResult, Error> {
342    let model = resolve_model(&model, ProviderRegistry::embedding_model)?;
343    let identity = ModelIdentity::new(model.provider().clone(), model.model_id().clone());
344    let span = spans::modality_span("embed", &identity);
345    base.run(|base, token| {
346        async move { run_calls(model, identity, values, &base, max_parallel_calls, token).await }
347            .instrument(span)
348    })
349    .await
350}
351
352async fn run_calls(
353    model: Arc<dyn DynEmbeddingModel>,
354    identity: ModelIdentity,
355    values: Vec<String>,
356    base: &ModalityOptions,
357    max_parallel_calls: Option<usize>,
358    cancellation: CancellationToken,
359) -> Result<EmbedManyResult, Error> {
360    let telemetry = TelemetryDispatcher::new(base.telemetry.clone());
361    let call_id = default_id_generator().generate();
362    let headers = base.request_headers();
363
364    let max_embeddings = match model.max_embeddings_per_call() {
365        Some(0) => {
366            return Err(Error::invalid_argument(
367                "max_embeddings_per_call",
368                "must be greater than 0",
369            ));
370        }
371        Some(limit) => limit,
372        None => usize::MAX,
373    };
374    let max_bytes = match model.max_input_bytes_per_call() {
375        Some(0) => {
376            return Err(Error::invalid_argument(
377                "max_input_bytes_per_call",
378                "must be greater than 0",
379            ));
380        }
381        Some(limit) => limit,
382        None => usize::MAX,
383    };
384    let chunks = split_by_limits(&values, max_embeddings, max_bytes);
385    let parallel = if model.supports_parallel_calls() {
386        max_parallel_calls.unwrap_or(usize::MAX).max(1)
387    } else {
388        1
389    };
390    let make_call = |values: Vec<String>| ChunkCall {
391        model: Arc::clone(&model),
392        identity: identity.clone(),
393        values,
394        headers: headers.clone(),
395        provider_options: base.provider_options.clone(),
396        retry_policy: base.retry_policy.clone(),
397        cancellation: cancellation.clone(),
398        telemetry: telemetry.clone(),
399        call_id: call_id.clone(),
400    };
401
402    let mut results: Vec<Option<ModelEmbedResult>> = (0..chunks.len()).map(|_| None).collect();
403    let indexed: Vec<(usize, Vec<String>)> = chunks.into_iter().enumerate().collect();
404    for window in indexed.chunks(parallel) {
405        if let [(index, values)] = window {
406            let result = make_call(values.clone()).run().await?;
407            if let Some(slot) = results.get_mut(*index) {
408                *slot = Some(result);
409            }
410            continue;
411        }
412        let mut tasks: JoinSet<(usize, Result<ModelEmbedResult, Error>)> = JoinSet::new();
413        for (index, values) in window {
414            let call = make_call(values.clone());
415            let index = *index;
416            tasks.spawn(async move { (index, call.run().await) });
417        }
418        while let Some(joined) = tasks.join_next().await {
419            let (index, result) = joined
420                .map_err(|error| Error::message(format!("embedding task failed: {error}")))?;
421            if let Some(slot) = results.get_mut(index) {
422                *slot = Some(result?);
423            }
424        }
425    }
426
427    let mut embeddings: Vec<Embedding> = Vec::with_capacity(values.len());
428    let mut warnings: Vec<Warning> = Vec::new();
429    let mut responses: Vec<ResponseMetadata> = Vec::new();
430    let mut tokens: Option<u64> = Some(0);
431    let mut provider_metadata: Option<ProviderMetadata> = None;
432    for result in results.into_iter().flatten() {
433        embeddings.extend(result.embeddings);
434        warnings.extend(result.warnings);
435        responses.push(result.response);
436        tokens = match (tokens, result.usage) {
437            (Some(total), Some(usage)) => Some(total.saturating_add(usage.tokens)),
438            _ => None,
439        };
440        accumulate_provider_metadata(&mut provider_metadata, result.provider_metadata.as_ref());
441    }
442    if embeddings.len() != values.len() {
443        return Err(invalid_count(values.len(), embeddings.len()));
444    }
445    Ok(EmbedManyResult {
446        values,
447        embeddings,
448        usage: EmbeddingUsage { tokens },
449        warnings,
450        responses,
451        provider_metadata,
452    })
453}