use std::sync::Arc;
use std::{error::Error as StdError, fmt};
pub use fastembed::EmbeddingModel as FastembedModel;
#[cfg(feature = "hf-hub")]
use fastembed::InitOptions;
use fastembed::{InitOptionsUserDefined, TextEmbedding, UserDefinedEmbeddingModel};
use rig_core::driver::{Exchange, Local, Model, Opened, Opening, Step, Transport};
use rig_core::embeddings;
use rig_core::error::ProviderError;
use rig_core::operation::Embedding;
use rig_core::wire::Capabilities;
#[derive(Debug, Clone)]
pub enum FastembedError {
UnknownModel(FastembedModel),
Initialization(String),
}
impl fmt::Display for FastembedError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
FastembedError::UnknownModel(model) => {
write!(
f,
"Failed to resolve FastEmbed model metadata for {model:?}"
)
}
FastembedError::Initialization(message) => {
write!(f, "Failed to initialize FastEmbed model: {message}")
}
}
}
}
impl StdError for FastembedError {}
pub fn text_embeddings(
model: &FastembedModel,
ndims: Option<usize>,
) -> Result<Local<Embedding>, FastembedError> {
let ndims = match ndims {
Some(ndims) => ndims,
None => TextEmbedding::get_model_info(model)
.map(|info| info.dim)
.map_err(|_| FastembedError::UnknownModel(model.clone()))?,
};
Ok(Local::new("fastembed")
.with_id(format!("{model:?}"))
.with_capabilities(Capabilities::embedding(1024, ndims)))
}
#[derive(Clone)]
pub struct Fastembed {
embedder: Arc<TextEmbedding>,
}
impl Fastembed {
#[cfg(feature = "hf-hub")]
pub fn load(model: &FastembedModel) -> Result<Self, FastembedError> {
let embedder = TextEmbedding::try_new(
InitOptions::new(model.to_owned()).with_show_download_progress(true),
)
.map_err(|err| FastembedError::Initialization(err.to_string()))?;
Ok(Self {
embedder: Arc::new(embedder),
})
}
pub fn embedding(
&self,
model: &FastembedModel,
ndims: Option<usize>,
) -> Result<Model<Local<Embedding>, Self>, FastembedError> {
Ok(Model::new(text_embeddings(model, ndims)?, self.clone()))
}
pub fn from_user_defined(
user_defined_model: UserDefinedEmbeddingModel,
) -> Result<Self, FastembedError> {
let embedder = TextEmbedding::try_new_from_user_defined(
user_defined_model,
InitOptionsUserDefined::default(),
)
.map_err(|err| FastembedError::Initialization(err.to_string()))?;
Ok(Self {
embedder: Arc::new(embedder),
})
}
}
impl Transport<Local<Embedding>> for Fastembed {
fn send(&self, texts: Vec<String>, _exchange: Exchange) -> Opening<Step<Embedding>> {
let embedder = Arc::clone(&self.embedder);
Opening::new(async move {
let embedded = embedder
.embed(texts.iter().map(String::as_str).collect(), None)
.map(|vectors| {
let embeddings = texts
.into_iter()
.zip(vectors)
.map(|(document, vector)| embeddings::Embedding {
document,
vec: vector.into_iter().map(f64::from).collect(),
})
.collect();
embeddings::EmbeddingResponse::new(embeddings)
})
.map_err(|err| ProviderError::Provider(err.to_string()));
Ok(Opened::new(futures::stream::iter(
[embedded.map(Step::End)],
)))
})
}
}