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, |_| {
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}