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, |attempt| {
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 = format!("{call_id}/attempt/{attempt}");
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                match &result {
308                    Ok(result) => telemetry.on_embed_end(&EmbedEndEvent {
309                        call_id: call_id.clone(),
310                        embedding_count: result.embeddings.len(),
311                        tokens: result.usage.map(|usage| usage.tokens),
312                        duration: started.elapsed(),
313                    }),
314                    Err(error) => telemetry.on_error(&ErrorEvent {
315                        call_id: &call_id,
316                        error,
317                        phase: ErrorPhase::ModelCall,
318                    }),
319                }
320                result
321            }
322        })
323        .await;
324        let result = outcome?;
325        if result.embeddings.len() != values.len() {
326            return Err(invalid_count(values.len(), result.embeddings.len()));
327        }
328        spans::log_warnings(&result.warnings, &identity);
329        Ok(result)
330    }
331}
332
333async fn run(
334    model: EmbeddingModelRef,
335    values: Vec<String>,
336    base: ModalityOptions,
337    max_parallel_calls: Option<usize>,
338) -> Result<EmbedManyResult, Error> {
339    let model = resolve_model(&model, ProviderRegistry::embedding_model)?;
340    let identity = ModelIdentity::new(model.provider().clone(), model.model_id().clone());
341    let span = spans::modality_span("embed", &identity);
342    base.run(|base, token| {
343        async move { run_calls(model, identity, values, &base, max_parallel_calls, token).await }
344            .instrument(span)
345    })
346    .await
347}
348
349async fn run_calls(
350    model: Arc<dyn DynEmbeddingModel>,
351    identity: ModelIdentity,
352    values: Vec<String>,
353    base: &ModalityOptions,
354    max_parallel_calls: Option<usize>,
355    cancellation: CancellationToken,
356) -> Result<EmbedManyResult, Error> {
357    let telemetry = TelemetryDispatcher::new(base.telemetry.clone());
358    let call_id = default_id_generator().generate();
359    let headers = base.request_headers();
360
361    let max_embeddings = match model.max_embeddings_per_call() {
362        Some(0) => {
363            return Err(Error::invalid_argument(
364                "max_embeddings_per_call",
365                "must be greater than 0",
366            ));
367        }
368        Some(limit) => limit,
369        None => usize::MAX,
370    };
371    let max_bytes = match model.max_input_bytes_per_call() {
372        Some(0) => {
373            return Err(Error::invalid_argument(
374                "max_input_bytes_per_call",
375                "must be greater than 0",
376            ));
377        }
378        Some(limit) => limit,
379        None => usize::MAX,
380    };
381    let chunks = split_by_limits(&values, max_embeddings, max_bytes);
382    let parallel = if model.supports_parallel_calls() {
383        max_parallel_calls.unwrap_or(usize::MAX).max(1)
384    } else {
385        1
386    };
387    let make_call = |index: usize, values: Vec<String>| ChunkCall {
388        model: Arc::clone(&model),
389        identity: identity.clone(),
390        values,
391        headers: headers.clone(),
392        provider_options: base.provider_options.clone(),
393        retry_policy: base.retry_policy.clone(),
394        cancellation: cancellation.clone(),
395        telemetry: telemetry.clone(),
396        call_id: format!("{call_id}/chunk/{index}"),
397    };
398
399    let mut results: Vec<Option<ModelEmbedResult>> = (0..chunks.len()).map(|_| None).collect();
400    let indexed: Vec<(usize, Vec<String>)> = chunks.into_iter().enumerate().collect();
401    for window in indexed.chunks(parallel) {
402        if let [(index, values)] = window {
403            let result = make_call(*index, values.clone()).run().await?;
404            if let Some(slot) = results.get_mut(*index) {
405                *slot = Some(result);
406            }
407            continue;
408        }
409        let mut tasks: JoinSet<(usize, Result<ModelEmbedResult, Error>)> = JoinSet::new();
410        for (index, values) in window {
411            let call = make_call(*index, values.clone());
412            let index = *index;
413            tasks.spawn(async move { (index, call.run().await) });
414        }
415        while let Some(joined) = tasks.join_next().await {
416            let (index, result) = joined
417                .map_err(|error| Error::message(format!("embedding task failed: {error}")))?;
418            if let Some(slot) = results.get_mut(index) {
419                *slot = Some(result?);
420            }
421        }
422    }
423
424    let mut embeddings: Vec<Embedding> = Vec::with_capacity(values.len());
425    let mut warnings: Vec<Warning> = Vec::new();
426    let mut responses: Vec<ResponseMetadata> = Vec::new();
427    let mut tokens: Option<u64> = Some(0);
428    let mut provider_metadata: Option<ProviderMetadata> = None;
429    for result in results.into_iter().flatten() {
430        embeddings.extend(result.embeddings);
431        warnings.extend(result.warnings);
432        responses.push(result.response);
433        tokens = match (tokens, result.usage) {
434            (Some(total), Some(usage)) => Some(total.saturating_add(usage.tokens)),
435            _ => None,
436        };
437        accumulate_provider_metadata(&mut provider_metadata, result.provider_metadata.as_ref());
438    }
439    if embeddings.len() != values.len() {
440        return Err(invalid_count(values.len(), embeddings.len()));
441    }
442    Ok(EmbedManyResult {
443        values,
444        embeddings,
445        usage: EmbeddingUsage { tokens },
446        warnings,
447        responses,
448        provider_metadata,
449    })
450}