use std::fmt;
use std::sync::Arc;
use ferrin_spec::BoxFuture;
use ferrin_spec::DynImageModel;
use ferrin_spec::ImageModel;
use ferrin_spec::ModelId;
use ferrin_spec::ProviderId;
use ferrin_spec::error::ProviderError;
use ferrin_spec::image_model::ImageOptions;
use ferrin_spec::image_model::ImageResult;
#[derive(Clone, Copy)]
pub struct ImageMiddlewareContext<'a> {
pub model: &'a dyn DynImageModel,
}
impl fmt::Debug for ImageMiddlewareContext<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ImageMiddlewareContext")
.field("provider", self.model.provider())
.field("model_id", self.model.model_id())
.finish()
}
}
pub type ImageGenerateNext<'a> =
Box<dyn FnOnce(ImageOptions) -> BoxFuture<'a, Result<ImageResult, ProviderError>> + Send + 'a>;
pub trait ImageModelMiddleware: Send + Sync + 'static {
fn transform_params<'a>(
&'a self,
options: ImageOptions,
_ctx: ImageMiddlewareContext<'a>,
) -> BoxFuture<'a, Result<ImageOptions, ProviderError>> {
Box::pin(async move { Ok(options) })
}
fn wrap_generate<'a>(
&'a self,
options: ImageOptions,
next: ImageGenerateNext<'a>,
_ctx: ImageMiddlewareContext<'a>,
) -> BoxFuture<'a, Result<ImageResult, ProviderError>> {
next(options)
}
fn override_provider(&self, _model: &dyn DynImageModel) -> Option<ProviderId> {
None
}
fn override_model_id(&self, _model: &dyn DynImageModel) -> Option<ModelId> {
None
}
fn max_images_per_call(&self, model: &dyn DynImageModel) -> Option<usize> {
model.max_images_per_call()
}
}
#[must_use]
pub fn wrap_image_model(
model: Arc<dyn DynImageModel>,
middleware: impl IntoIterator<Item = Arc<dyn ImageModelMiddleware>, IntoIter: DoubleEndedIterator>,
) -> Arc<dyn DynImageModel> {
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(WrappedImageModel {
inner,
layer,
provider,
model_id,
})
})
}
struct WrappedImageModel {
inner: Arc<dyn DynImageModel>,
layer: Arc<dyn ImageModelMiddleware>,
provider: ProviderId,
model_id: ModelId,
}
impl fmt::Debug for WrappedImageModel {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WrappedImageModel")
.field("provider", &self.provider)
.field("model_id", &self.model_id)
.finish_non_exhaustive()
}
}
impl ImageModel for WrappedImageModel {
fn provider(&self) -> &ProviderId {
&self.provider
}
fn model_id(&self) -> &ModelId {
&self.model_id
}
fn max_images_per_call(&self) -> Option<usize> {
self.layer.max_images_per_call(self.inner.as_ref())
}
async fn do_generate(&self, options: ImageOptions) -> Result<ImageResult, ProviderError> {
let ctx = ImageMiddlewareContext {
model: self.inner.as_ref(),
};
let options = self.layer.transform_params(options, ctx).await?;
let inner = &self.inner;
self.layer
.wrap_generate(
options,
Box::new(move |options| inner.do_generate(options)),
ctx,
)
.await
}
}