Skip to main content

rig_fastembed/
lib.rs

1//! Local embedding model integration backed by `fastembed`.
2//!
3//! A loaded `fastembed` model is the [`Fastembed`] transport, the runtime
4//! behind a local embedding wire ([`text_embeddings`]) that embeds in the
5//! calling process. The default feature set enables Hugging Face model
6//! downloads and ONNX Runtime binary downloads.
7//!
8//! ```no_run
9//! use rig_core::Model;
10//! use rig_fastembed::{Fastembed, FastembedModel, text_embeddings};
11//!
12//! # fn run() -> Result<(), rig_fastembed::FastembedError> {
13//! let model = Fastembed::load(&FastembedModel::AllMiniLML6V2Q)?.embedding(&FastembedModel::AllMiniLML6V2Q, None)?;
14//! # let _ = model;
15//! # Ok(())
16//! # }
17//! ```
18//!
19//! `rig-fastembed` is native-only and does not target `wasm32-unknown-unknown`.
20//! The root `rig` facade re-exports this crate as `rig::fastembed` when one of
21//! its Fastembed features is enabled.
22
23use std::sync::Arc;
24use std::{error::Error as StdError, fmt};
25
26pub use fastembed::EmbeddingModel as FastembedModel;
27#[cfg(feature = "hf-hub")]
28use fastembed::InitOptions;
29use fastembed::{InitOptionsUserDefined, TextEmbedding, UserDefinedEmbeddingModel};
30use rig_core::driver::{Exchange, Local, Model, Opened, Opening, Step, Transport};
31use rig_core::embeddings;
32use rig_core::error::ProviderError;
33use rig_core::operation::Embedding;
34use rig_core::wire::Capabilities;
35
36/// Errors raised while resolving or initializing a Fastembed model.
37#[derive(Debug, Clone)]
38pub enum FastembedError {
39    /// `fastembed` has no metadata for the requested model.
40    UnknownModel(FastembedModel),
41    /// The model failed to load, download, or initialize.
42    Initialization(String),
43}
44
45impl fmt::Display for FastembedError {
46    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
47        match self {
48            FastembedError::UnknownModel(model) => {
49                write!(
50                    f,
51                    "Failed to resolve FastEmbed model metadata for {model:?}"
52                )
53            }
54            FastembedError::Initialization(message) => {
55                write!(f, "Failed to initialize FastEmbed model: {message}")
56            }
57        }
58    }
59}
60
61impl StdError for FastembedError {}
62
63/// The local embedding wire of `model` at `ndims` dimensions, named
64/// `fastembed` and addressing the model by its `fastembed` name. `None`
65/// takes the width from the model metadata, which errors for models
66/// `fastembed` does not know.
67pub fn text_embeddings(
68    model: &FastembedModel,
69    ndims: Option<usize>,
70) -> Result<Local<Embedding>, FastembedError> {
71    let ndims = match ndims {
72        Some(ndims) => ndims,
73        None => TextEmbedding::get_model_info(model)
74            .map(|info| info.dim)
75            .map_err(|_| FastembedError::UnknownModel(model.clone()))?,
76    };
77    Ok(Local::new("fastembed")
78        .with_id(format!("{model:?}"))
79        .with_capabilities(Capabilities::embedding(1024, ndims)))
80}
81
82/// A loaded Fastembed model: the transport that embeds in the calling
83/// process. Clones share the loaded model. Pair it with the
84/// [`text_embeddings`] wire of the model it loaded: the wire names the model
85/// and width that spans and capabilities report, and the transport embeds
86/// with whatever it loaded. In-process execution reports no raw payload,
87/// usage, or request id.
88#[derive(Clone)]
89pub struct Fastembed {
90    embedder: Arc<TextEmbedding>,
91}
92
93impl Fastembed {
94    /// Loads `model`, downloading it when necessary and reporting download
95    /// progress on standard output.
96    #[cfg(feature = "hf-hub")]
97    pub fn load(model: &FastembedModel) -> Result<Self, FastembedError> {
98        let embedder = TextEmbedding::try_new(
99            InitOptions::new(model.to_owned()).with_show_download_progress(true),
100        )
101        .map_err(|err| FastembedError::Initialization(err.to_string()))?;
102        Ok(Self {
103            embedder: Arc::new(embedder),
104        })
105    }
106
107    /// The embedding model of `model` at `ndims` dimensions, on this
108    /// runtime: the [`text_embeddings`] wire, which fails for a model
109    /// `fastembed` does not know when `ndims` is `None`. `model` names what
110    /// spans and capabilities report; this runtime embeds with whatever it
111    /// loaded.
112    pub fn embedding(
113        &self,
114        model: &FastembedModel,
115        ndims: Option<usize>,
116    ) -> Result<Model<Local<Embedding>, Self>, FastembedError> {
117        Ok(Model::new(text_embeddings(model, ndims)?, self.clone()))
118    }
119
120    /// Loads a caller-supplied ONNX model.
121    pub fn from_user_defined(
122        user_defined_model: UserDefinedEmbeddingModel,
123    ) -> Result<Self, FastembedError> {
124        let embedder = TextEmbedding::try_new_from_user_defined(
125            user_defined_model,
126            InitOptionsUserDefined::default(),
127        )
128        .map_err(|err| FastembedError::Initialization(err.to_string()))?;
129        Ok(Self {
130            embedder: Arc::new(embedder),
131        })
132    }
133}
134
135impl Transport<Local<Embedding>> for Fastembed {
136    fn send(&self, texts: Vec<String>, _exchange: Exchange) -> Opening<Step<Embedding>> {
137        let embedder = Arc::clone(&self.embedder);
138        Opening::new(async move {
139            let embedded = embedder
140                .embed(texts.iter().map(String::as_str).collect(), None)
141                .map(|vectors| {
142                    let embeddings = texts
143                        .into_iter()
144                        .zip(vectors)
145                        .map(|(document, vector)| embeddings::Embedding {
146                            document,
147                            vec: vector.into_iter().map(f64::from).collect(),
148                        })
149                        .collect();
150                    embeddings::EmbeddingResponse::new(embeddings)
151                })
152                .map_err(|err| ProviderError::Provider(err.to_string()));
153            // A failed embed fails the reply, as a transport failure does.
154            Ok(Opened::new(futures::stream::iter(
155                [embedded.map(Step::End)],
156            )))
157        })
158    }
159}