ferrin_core/middleware/
embedding.rs1use 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#[derive(Clone, Copy)]
23pub struct EmbeddingMiddlewareContext<'a> {
24 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
37pub type EmbedNext<'a> =
39 Box<dyn FnOnce(EmbedOptions) -> BoxFuture<'a, Result<EmbedResult, ProviderError>> + Send + 'a>;
40
41pub trait EmbeddingModelMiddleware: Send + Sync + 'static {
47 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 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 fn override_provider(&self, _model: &dyn DynEmbeddingModel) -> Option<ProviderId> {
68 None
69 }
70
71 fn override_model_id(&self, _model: &dyn DynEmbeddingModel) -> Option<ModelId> {
73 None
74 }
75
76 fn max_embeddings_per_call(&self, model: &dyn DynEmbeddingModel) -> Option<usize> {
78 model.max_embeddings_per_call()
79 }
80
81 fn max_input_bytes_per_call(&self, model: &dyn DynEmbeddingModel) -> Option<usize> {
84 model.max_input_bytes_per_call()
85 }
86
87 fn supports_parallel_calls(&self, model: &dyn DynEmbeddingModel) -> bool {
90 model.supports_parallel_calls()
91 }
92}
93
94#[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}