use std::future::IntoFuture;
use std::sync::Arc;
use std::time::Instant;
use ferrin_spec::BoxFuture;
use ferrin_spec::DynEmbeddingModel;
use ferrin_spec::EmbeddingModelRef;
use ferrin_spec::Headers;
use ferrin_spec::ProviderMetadata;
use ferrin_spec::ProviderOptions;
use ferrin_spec::ResponseMetadata;
use ferrin_spec::Warning;
use ferrin_spec::embedding_model::EmbedOptions;
use ferrin_spec::embedding_model::EmbedResult as ModelEmbedResult;
pub use ferrin_spec::embedding_model::Embedding;
use ferrin_spec::error::InvalidResponseDataError;
use ferrin_spec::error::ProviderError;
use serde_json::json;
use tokio::task::JoinSet;
use tokio_util::sync::CancellationToken;
use tracing::Instrument;
use crate::error::Error;
use crate::hooks::Hooks;
use crate::ids::default_id_generator;
use crate::modality::ModalityOptions;
use crate::modality::impl_modality_builder;
pub use crate::modality_hooks::EmbedCallEndEvent;
pub use crate::modality_hooks::EmbedCallStartEvent;
pub use crate::modality_hooks::EmbeddingInput;
pub use crate::modality_hooks::EmbeddingOutput;
pub use crate::modality_hooks::EmbeddingResponse;
use crate::modality_hooks::ModalityHooks;
use crate::modality_hooks::impl_modality_hooks;
use crate::modality_metadata::accumulate_embedding_metadata;
use crate::registry::ProviderRegistry;
use crate::registry::default::resolve_model;
use crate::retry::RetryPolicy;
use crate::retry::retry;
use crate::telemetry::EmbedEndEvent;
use crate::telemetry::EmbedStartEvent;
use crate::telemetry::ErrorEvent;
use crate::telemetry::ErrorPhase;
use crate::telemetry::ModelIdentity;
use crate::telemetry::dispatcher::TelemetryDispatcher;
use crate::telemetry::spans;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct EmbeddingUsage {
pub tokens: Option<u64>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct EmbedResult {
pub value: String,
pub embedding: Embedding,
pub usage: EmbeddingUsage,
pub warnings: Vec<Warning>,
pub response: ResponseMetadata,
pub provider_metadata: Option<ProviderMetadata>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct EmbedManyResult {
pub values: Vec<String>,
pub embeddings: Vec<Embedding>,
pub usage: EmbeddingUsage,
pub warnings: Vec<Warning>,
pub responses: Vec<ResponseMetadata>,
pub provider_metadata: Option<ProviderMetadata>,
}
#[must_use]
pub fn embed(model: impl Into<EmbeddingModelRef>, value: impl Into<String>) -> Embed {
Embed {
model: model.into(),
value: value.into(),
base: ModalityOptions::default(),
hooks: ModalityHooks::default(),
}
}
#[derive(Debug)]
pub struct Embed {
model: EmbeddingModelRef,
value: String,
base: ModalityOptions,
hooks: ModalityHooks<EmbedCallStartEvent, EmbedCallEndEvent>,
}
impl_modality_builder!(Embed);
impl_modality_hooks!(Embed, EmbedCallStartEvent, EmbedCallEndEvent);
impl IntoFuture for Embed {
type Output = Result<EmbedResult, Error>;
type IntoFuture = BoxFuture<'static, Self::Output>;
fn into_future(self) -> Self::IntoFuture {
Box::pin(async move {
let many = run(
self.model,
EmbeddingInput::Single(self.value),
self.base,
Some(1),
self.hooks,
)
.await?;
let EmbedManyResult {
values,
embeddings,
usage,
warnings,
responses,
provider_metadata,
} = many;
let (Some(value), Some(embedding), Some(response)) = (
values.into_iter().next(),
embeddings.into_iter().next(),
responses.into_iter().next(),
) else {
return Err(invalid_count(1, 0));
};
Ok(EmbedResult {
value,
embedding,
usage,
warnings,
response,
provider_metadata,
})
})
}
}
#[must_use]
pub fn embed_many(
model: impl Into<EmbeddingModelRef>,
values: impl IntoIterator<Item = impl Into<String>>,
) -> EmbedMany {
EmbedMany {
model: model.into(),
values: values.into_iter().map(Into::into).collect(),
max_parallel_calls: None,
base: ModalityOptions::default(),
hooks: ModalityHooks::default(),
}
}
#[derive(Debug)]
pub struct EmbedMany {
model: EmbeddingModelRef,
values: Vec<String>,
max_parallel_calls: Option<usize>,
base: ModalityOptions,
hooks: ModalityHooks<EmbedCallStartEvent, EmbedCallEndEvent>,
}
impl EmbedMany {
#[must_use]
pub fn max_parallel_calls(mut self, max_parallel_calls: usize) -> Self {
self.max_parallel_calls = Some(max_parallel_calls.max(1));
self
}
}
impl_modality_builder!(EmbedMany);
impl_modality_hooks!(EmbedMany, EmbedCallStartEvent, EmbedCallEndEvent);
impl IntoFuture for EmbedMany {
type Output = Result<EmbedManyResult, Error>;
type IntoFuture = BoxFuture<'static, Self::Output>;
fn into_future(self) -> Self::IntoFuture {
Box::pin(run(
self.model,
EmbeddingInput::Many(self.values),
self.base,
self.max_parallel_calls,
self.hooks,
))
}
}
pub fn cosine_similarity(a: &[f64], b: &[f64]) -> Result<f64, Error> {
if a.len() != b.len() {
return Err(Error::invalid_argument(
"vectors",
format!(
"vectors must have the same length (got {} and {})",
a.len(),
b.len()
),
));
}
let dot: f64 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let norm_a = a.iter().map(|x| x * x).sum::<f64>().sqrt();
let norm_b = b.iter().map(|y| y * y).sum::<f64>().sqrt();
if norm_a == 0.0 || norm_b == 0.0 {
return Ok(0.0);
}
Ok(dot / (norm_a * norm_b))
}
pub(crate) fn split_by_limits(
values: &[String],
max_embeddings: usize,
max_bytes: usize,
) -> Vec<Vec<String>> {
let mut chunks: Vec<Vec<String>> = Vec::new();
let mut current: Vec<String> = Vec::new();
let mut current_bytes = 0usize;
for value in values {
let bytes = value.len();
if !current.is_empty()
&& (current.len() >= max_embeddings || current_bytes.saturating_add(bytes) > max_bytes)
{
chunks.push(std::mem::take(&mut current));
current_bytes = 0;
}
current.push(value.clone());
current_bytes = current_bytes.saturating_add(bytes);
}
if !current.is_empty() {
chunks.push(current);
}
chunks
}
fn invalid_count(expected: usize, received: usize) -> Error {
Error::from(ProviderError::InvalidResponseData(Box::new(
InvalidResponseDataError::new(
format!("expected {expected} embeddings, received {received}"),
json!({ "expected": expected, "received": received }),
),
)))
}
struct ChunkCall {
model: Arc<dyn DynEmbeddingModel>,
identity: ModelIdentity,
values: Vec<String>,
headers: Headers,
provider_options: ProviderOptions,
retry_policy: RetryPolicy,
cancellation: CancellationToken,
telemetry: TelemetryDispatcher,
call_id: String,
}
impl ChunkCall {
async fn run(self) -> Result<ModelEmbedResult, Error> {
let Self {
model,
identity,
values,
headers,
provider_options,
retry_policy,
cancellation,
telemetry,
call_id,
} = self;
let outcome = retry(&retry_policy, &cancellation, |attempt| {
let values = values.clone();
let model = &model;
let identity = &identity;
let headers = &headers;
let provider_options = &provider_options;
let cancellation = &cancellation;
let telemetry = &telemetry;
let call_id = format!("{call_id}/attempt/{attempt}");
async move {
let started = Instant::now();
telemetry
.on_embed_start(&EmbedStartEvent {
call_id: call_id.clone(),
model: identity.clone(),
value_count: values.len(),
values: telemetry.record_inputs().then(|| values.clone()),
})
.await;
let result = model
.do_embed(EmbedOptions {
values,
headers: headers.clone(),
provider_options: provider_options.clone(),
cancellation: cancellation.child_token(),
})
.await
.map_err(Error::from);
match &result {
Ok(result) => {
telemetry
.on_embed_end(&EmbedEndEvent {
call_id: call_id.clone(),
embedding_count: result.embeddings.len(),
tokens: result.usage.map(|usage| usage.tokens),
duration: started.elapsed(),
})
.await
}
Err(error) => {
telemetry
.on_error(&ErrorEvent {
call_id: &call_id,
error,
phase: ErrorPhase::ModelCall,
})
.await
}
}
result
}
})
.await;
let result = outcome?;
if result.embeddings.len() != values.len() {
return Err(invalid_count(values.len(), result.embeddings.len()));
}
spans::log_warnings(&result.warnings, &identity);
Ok(result)
}
}
async fn run(
model: EmbeddingModelRef,
value: EmbeddingInput,
base: ModalityOptions,
max_parallel_calls: Option<usize>,
hooks: ModalityHooks<EmbedCallStartEvent, EmbedCallEndEvent>,
) -> Result<EmbedManyResult, Error> {
let model = resolve_model(&model, ProviderRegistry::embedding_model)?;
let identity = ModelIdentity::new(model.provider().clone(), model.model_id().clone());
let span = spans::modality_span("embed", &identity);
base.run(|base, token| {
async move {
let telemetry = TelemetryDispatcher::new(base.telemetry.clone());
let call_id = default_id_generator().generate();
let (operation_id, values) = match &value {
EmbeddingInput::Single(value) => ("ai.embed", vec![value.clone()]),
EmbeddingInput::Many(values) => ("ai.embedMany", values.clone()),
};
let start = Arc::new(EmbedCallStartEvent {
runtime_context: Some(hooks.runtime_context.clone()),
call_id: call_id.clone(),
operation_id,
model: identity.clone(),
value: Some(value.clone()),
max_retries: base.retry_policy.max_retries,
headers: base.request_headers(),
provider_options: base.provider_options.clone(),
});
tokio::join!(
Hooks::emit(&hooks.on_start, start.clone()),
telemetry.on_embed_operation_start(&start),
);
let result = run_calls(
model,
identity.clone(),
values,
&base,
max_parallel_calls,
token,
&call_id,
)
.await?;
let (embedding, response) = match &value {
EmbeddingInput::Single(_) => (
EmbeddingOutput::Single(
result
.embeddings
.first()
.cloned()
.ok_or_else(|| invalid_count(1, 0))?,
),
EmbeddingResponse::Single(Box::new(
result.responses.first().cloned().unwrap_or_default(),
)),
),
EmbeddingInput::Many(_) => (
EmbeddingOutput::Many(result.embeddings.clone()),
EmbeddingResponse::Many(result.responses.clone()),
),
};
let end = Arc::new(EmbedCallEndEvent {
runtime_context: Some(hooks.runtime_context),
call_id,
operation_id,
model: identity,
value: Some(value),
embedding: Some(embedding),
usage: result.usage,
warnings: result.warnings.clone(),
provider_metadata: result.provider_metadata.clone(),
response,
});
tokio::join!(
Hooks::emit(&hooks.on_end, end.clone()),
telemetry.on_embed_operation_end(&end),
);
Ok(result)
}
.instrument(span)
})
.await
}
async fn run_calls(
model: Arc<dyn DynEmbeddingModel>,
identity: ModelIdentity,
values: Vec<String>,
base: &ModalityOptions,
max_parallel_calls: Option<usize>,
cancellation: CancellationToken,
call_id: &str,
) -> Result<EmbedManyResult, Error> {
let telemetry = TelemetryDispatcher::new(base.telemetry.clone());
let headers = base.request_headers();
let max_embeddings = match model.max_embeddings_per_call() {
Some(0) => {
return Err(Error::invalid_argument(
"max_embeddings_per_call",
"must be greater than 0",
));
}
Some(limit) => limit,
None => usize::MAX,
};
let max_bytes = match model.max_input_bytes_per_call() {
Some(0) => {
return Err(Error::invalid_argument(
"max_input_bytes_per_call",
"must be greater than 0",
));
}
Some(limit) => limit,
None => usize::MAX,
};
let chunks = if model.max_embeddings_per_call().is_none()
&& model.max_input_bytes_per_call().is_none()
{
vec![values.clone()]
} else {
split_by_limits(&values, max_embeddings, max_bytes)
};
let parallel = if model.supports_parallel_calls() {
max_parallel_calls.unwrap_or(usize::MAX).max(1)
} else {
1
};
let make_call = |index: usize, values: Vec<String>| ChunkCall {
model: Arc::clone(&model),
identity: identity.clone(),
values,
headers: headers.clone(),
provider_options: base.provider_options.clone(),
retry_policy: base.retry_policy.clone(),
cancellation: cancellation.clone(),
telemetry: telemetry.clone(),
call_id: format!("{call_id}/chunk/{index}"),
};
let mut results: Vec<Option<ModelEmbedResult>> = (0..chunks.len()).map(|_| None).collect();
let indexed: Vec<(usize, Vec<String>)> = chunks.into_iter().enumerate().collect();
for window in indexed.chunks(parallel) {
if let [(index, values)] = window {
let result = make_call(*index, values.clone()).run().await?;
if let Some(slot) = results.get_mut(*index) {
*slot = Some(result);
}
continue;
}
let mut tasks: JoinSet<(usize, Result<ModelEmbedResult, Error>)> = JoinSet::new();
for (index, values) in window {
let call = make_call(*index, values.clone());
let index = *index;
tasks.spawn(async move { (index, call.run().await) });
}
while let Some(joined) = tasks.join_next().await {
let (index, result) = joined
.map_err(|error| Error::message(format!("embedding task failed: {error}")))?;
if let Some(slot) = results.get_mut(index) {
*slot = Some(result?);
}
}
}
let mut embeddings: Vec<Embedding> = Vec::with_capacity(values.len());
let mut warnings: Vec<Warning> = Vec::new();
let mut responses: Vec<ResponseMetadata> = Vec::new();
let mut tokens: Option<u64> = Some(0);
let mut provider_metadata: Option<ProviderMetadata> = None;
for result in results.into_iter().flatten() {
embeddings.extend(result.embeddings);
warnings.extend(result.warnings);
responses.push(result.response);
tokens = match (tokens, result.usage) {
(Some(total), Some(usage)) => Some(total.saturating_add(usage.tokens)),
_ => None,
};
accumulate_embedding_metadata(&mut provider_metadata, result.provider_metadata.as_ref());
}
if embeddings.len() != values.len() {
return Err(invalid_count(values.len(), embeddings.len()));
}
Ok(EmbedManyResult {
values,
embeddings,
usage: EmbeddingUsage { tokens },
warnings,
responses,
provider_metadata,
})
}