1use 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#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
58pub struct EmbeddingUsage {
59 pub tokens: Option<u64>,
61}
62
63#[derive(Debug, Clone, PartialEq)]
65pub struct EmbedResult {
66 pub value: String,
68 pub embedding: Embedding,
70 pub usage: EmbeddingUsage,
72 pub warnings: Vec<Warning>,
74 pub response: ResponseMetadata,
76 pub provider_metadata: Option<ProviderMetadata>,
78}
79
80#[derive(Debug, Clone, PartialEq)]
82pub struct EmbedManyResult {
83 pub values: Vec<String>,
85 pub embeddings: Vec<Embedding>,
87 pub usage: EmbeddingUsage,
89 pub warnings: Vec<Warning>,
91 pub responses: Vec<ResponseMetadata>,
93 pub provider_metadata: Option<ProviderMetadata>,
95}
96
97#[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#[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#[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#[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 #[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
216pub 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
242pub(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
280struct 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}