Skip to main content

ferrin_core/middleware/
embedding.rs

1//! Embedding model middleware.
2//!
3//! Mirrors the language model middleware for [`EmbeddingModel`]: a layer may
4//! rewrite the embed options, wrap `do_embed`, and override the identity or
5//! batching limits reported by the wrapped model. Apply with
6//! [`wrap_embedding_model`]; the first middleware in the list is the
7//! outermost.
8
9use std::fmt;
10use std::sync::Arc;
11
12use ferrin_spec::BoxFuture;
13use ferrin_spec::DynEmbeddingModel;
14use ferrin_spec::EmbeddingModel;
15use ferrin_spec::ModelId;
16use ferrin_spec::ProviderId;
17use ferrin_spec::embedding_model::EmbedOptions;
18use ferrin_spec::embedding_model::EmbedResult;
19use ferrin_spec::error::ProviderError;
20
21/// The embedding model being wrapped.
22#[derive(Clone, Copy)]
23pub struct EmbeddingMiddlewareContext<'a> {
24    /// The wrapped (inner) model.
25    pub model: &'a dyn DynEmbeddingModel,
26}
27
28impl fmt::Debug for EmbeddingMiddlewareContext<'_> {
29    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
30        f.debug_struct("EmbeddingMiddlewareContext")
31            .field("provider", self.model.provider())
32            .field("model_id", self.model.model_id())
33            .finish()
34    }
35}
36
37/// Continuation of a wrapped `do_embed`.
38pub type EmbedNext<'a> =
39    Box<dyn FnOnce(EmbedOptions) -> BoxFuture<'a, Result<EmbedResult, ProviderError>> + Send + 'a>;
40
41/// Intercepts embedding model calls. Every method has a pass-through default.
42///
43/// The limit methods receive the wrapped model and return the value the
44/// wrapper reports; their defaults forward the inner model's values, so a
45/// layer overrides a limit by returning something else.
46pub trait EmbeddingModelMiddleware: Send + Sync + 'static {
47    /// Rewrites the embed options before the call.
48    fn transform_params<'a>(
49        &'a self,
50        options: EmbedOptions,
51        _ctx: EmbeddingMiddlewareContext<'a>,
52    ) -> BoxFuture<'a, Result<EmbedOptions, ProviderError>> {
53        Box::pin(async move { Ok(options) })
54    }
55
56    /// Wraps `do_embed`.
57    fn wrap_embed<'a>(
58        &'a self,
59        options: EmbedOptions,
60        next: EmbedNext<'a>,
61        _ctx: EmbeddingMiddlewareContext<'a>,
62    ) -> BoxFuture<'a, Result<EmbedResult, ProviderError>> {
63        next(options)
64    }
65
66    /// Overrides the provider id reported by the wrapped model.
67    fn override_provider(&self, _model: &dyn DynEmbeddingModel) -> Option<ProviderId> {
68        None
69    }
70
71    /// Overrides the model id reported by the wrapped model.
72    fn override_model_id(&self, _model: &dyn DynEmbeddingModel) -> Option<ModelId> {
73        None
74    }
75
76    /// The batch limit reported by the wrapper; defaults to the inner value.
77    fn max_embeddings_per_call(&self, model: &dyn DynEmbeddingModel) -> Option<usize> {
78        model.max_embeddings_per_call()
79    }
80
81    /// The input byte limit reported by the wrapper; defaults to the inner
82    /// value.
83    fn max_input_bytes_per_call(&self, model: &dyn DynEmbeddingModel) -> Option<usize> {
84        model.max_input_bytes_per_call()
85    }
86
87    /// Whether the wrapper allows concurrent calls; defaults to the inner
88    /// value.
89    fn supports_parallel_calls(&self, model: &dyn DynEmbeddingModel) -> bool {
90        model.supports_parallel_calls()
91    }
92}
93
94/// Wraps `model` with `middleware`; the first entry becomes the outermost
95/// layer. An empty list returns `model` unchanged.
96#[must_use]
97pub fn wrap_embedding_model(
98    model: Arc<dyn DynEmbeddingModel>,
99    middleware: impl IntoIterator<
100        Item = Arc<dyn EmbeddingModelMiddleware>,
101        IntoIter: DoubleEndedIterator,
102    >,
103) -> Arc<dyn DynEmbeddingModel> {
104    middleware.into_iter().rev().fold(model, |inner, layer| {
105        let provider = layer
106            .override_provider(inner.as_ref())
107            .unwrap_or_else(|| inner.provider().clone());
108        let model_id = layer
109            .override_model_id(inner.as_ref())
110            .unwrap_or_else(|| inner.model_id().clone());
111        Arc::new(WrappedEmbeddingModel {
112            inner,
113            layer,
114            provider,
115            model_id,
116        })
117    })
118}
119
120struct WrappedEmbeddingModel {
121    inner: Arc<dyn DynEmbeddingModel>,
122    layer: Arc<dyn EmbeddingModelMiddleware>,
123    provider: ProviderId,
124    model_id: ModelId,
125}
126
127impl fmt::Debug for WrappedEmbeddingModel {
128    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
129        f.debug_struct("WrappedEmbeddingModel")
130            .field("provider", &self.provider)
131            .field("model_id", &self.model_id)
132            .finish_non_exhaustive()
133    }
134}
135
136impl EmbeddingModel for WrappedEmbeddingModel {
137    fn provider(&self) -> &ProviderId {
138        &self.provider
139    }
140
141    fn model_id(&self) -> &ModelId {
142        &self.model_id
143    }
144
145    fn max_embeddings_per_call(&self) -> Option<usize> {
146        self.layer.max_embeddings_per_call(self.inner.as_ref())
147    }
148
149    fn max_input_bytes_per_call(&self) -> Option<usize> {
150        self.layer.max_input_bytes_per_call(self.inner.as_ref())
151    }
152
153    fn supports_parallel_calls(&self) -> bool {
154        self.layer.supports_parallel_calls(self.inner.as_ref())
155    }
156
157    async fn do_embed(&self, options: EmbedOptions) -> Result<EmbedResult, ProviderError> {
158        let ctx = EmbeddingMiddlewareContext {
159            model: self.inner.as_ref(),
160        };
161        let options = self.layer.transform_params(options, ctx).await?;
162        let inner = &self.inner;
163        self.layer
164            .wrap_embed(
165                options,
166                Box::new(move |options| inner.do_embed(options)),
167                ctx,
168            )
169            .await
170    }
171}