use std::fmt;
use std::sync::Arc;
use ferrin_spec::BoxFuture;
use ferrin_spec::DynEmbeddingModel;
use ferrin_spec::EmbeddingModel;
use ferrin_spec::ModelId;
use ferrin_spec::ProviderId;
use ferrin_spec::embedding_model::EmbedOptions;
use ferrin_spec::embedding_model::EmbedResult;
use ferrin_spec::error::ProviderError;
#[derive(Clone, Copy)]
pub struct EmbeddingMiddlewareContext<'a> {
pub model: &'a dyn DynEmbeddingModel,
}
impl fmt::Debug for EmbeddingMiddlewareContext<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("EmbeddingMiddlewareContext")
.field("provider", self.model.provider())
.field("model_id", self.model.model_id())
.finish()
}
}
pub type EmbedNext<'a> =
Box<dyn FnOnce(EmbedOptions) -> BoxFuture<'a, Result<EmbedResult, ProviderError>> + Send + 'a>;
pub trait EmbeddingModelMiddleware: Send + Sync + 'static {
fn transform_params<'a>(
&'a self,
options: EmbedOptions,
_ctx: EmbeddingMiddlewareContext<'a>,
) -> BoxFuture<'a, Result<EmbedOptions, ProviderError>> {
Box::pin(async move { Ok(options) })
}
fn wrap_embed<'a>(
&'a self,
options: EmbedOptions,
next: EmbedNext<'a>,
_ctx: EmbeddingMiddlewareContext<'a>,
) -> BoxFuture<'a, Result<EmbedResult, ProviderError>> {
next(options)
}
fn override_provider(&self, _model: &dyn DynEmbeddingModel) -> Option<ProviderId> {
None
}
fn override_model_id(&self, _model: &dyn DynEmbeddingModel) -> Option<ModelId> {
None
}
fn max_embeddings_per_call(&self, model: &dyn DynEmbeddingModel) -> Option<usize> {
model.max_embeddings_per_call()
}
fn max_input_bytes_per_call(&self, model: &dyn DynEmbeddingModel) -> Option<usize> {
model.max_input_bytes_per_call()
}
fn supports_parallel_calls(&self, model: &dyn DynEmbeddingModel) -> bool {
model.supports_parallel_calls()
}
}
#[must_use]
pub fn wrap_embedding_model(
model: Arc<dyn DynEmbeddingModel>,
middleware: impl IntoIterator<
Item = Arc<dyn EmbeddingModelMiddleware>,
IntoIter: DoubleEndedIterator,
>,
) -> Arc<dyn DynEmbeddingModel> {
middleware.into_iter().rev().fold(model, |inner, layer| {
let provider = layer
.override_provider(inner.as_ref())
.unwrap_or_else(|| inner.provider().clone());
let model_id = layer
.override_model_id(inner.as_ref())
.unwrap_or_else(|| inner.model_id().clone());
Arc::new(WrappedEmbeddingModel {
inner,
layer,
provider,
model_id,
})
})
}
struct WrappedEmbeddingModel {
inner: Arc<dyn DynEmbeddingModel>,
layer: Arc<dyn EmbeddingModelMiddleware>,
provider: ProviderId,
model_id: ModelId,
}
impl fmt::Debug for WrappedEmbeddingModel {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WrappedEmbeddingModel")
.field("provider", &self.provider)
.field("model_id", &self.model_id)
.finish_non_exhaustive()
}
}
impl EmbeddingModel for WrappedEmbeddingModel {
fn provider(&self) -> &ProviderId {
&self.provider
}
fn model_id(&self) -> &ModelId {
&self.model_id
}
fn max_embeddings_per_call(&self) -> Option<usize> {
self.layer.max_embeddings_per_call(self.inner.as_ref())
}
fn max_input_bytes_per_call(&self) -> Option<usize> {
self.layer.max_input_bytes_per_call(self.inner.as_ref())
}
fn supports_parallel_calls(&self) -> bool {
self.layer.supports_parallel_calls(self.inner.as_ref())
}
async fn do_embed(&self, options: EmbedOptions) -> Result<EmbedResult, ProviderError> {
let ctx = EmbeddingMiddlewareContext {
model: self.inner.as_ref(),
};
let options = self.layer.transform_params(options, ctx).await?;
let inner = &self.inner;
self.layer
.wrap_embed(
options,
Box::new(move |options| inner.do_embed(options)),
ctx,
)
.await
}
}