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