1use 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#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
47pub struct EmbeddingUsage {
48 pub tokens: Option<u64>,
50}
51
52#[derive(Debug, Clone, PartialEq)]
54pub struct EmbedResult {
55 pub value: String,
57 pub embedding: Embedding,
59 pub usage: EmbeddingUsage,
61 pub warnings: Vec<Warning>,
63 pub response: ResponseMetadata,
65 pub provider_metadata: Option<ProviderMetadata>,
67}
68
69#[derive(Debug, Clone, PartialEq)]
71pub struct EmbedManyResult {
72 pub values: Vec<String>,
74 pub embeddings: Vec<Embedding>,
76 pub usage: EmbeddingUsage,
78 pub warnings: Vec<Warning>,
80 pub responses: Vec<ResponseMetadata>,
82 pub provider_metadata: Option<ProviderMetadata>,
84}
85
86#[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#[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#[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#[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 #[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
191pub 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
217pub(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
255struct 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}