Skip to main content

ferrin_core/middleware/
image.rs

1//! Image model middleware.
2//!
3//! Mirrors the language model middleware for [`ImageModel`]: a layer may
4//! rewrite the image options, wrap `do_generate`, and override the identity
5//! or per-call limit reported by the wrapped model. Apply with
6//! [`wrap_image_model`]; the first middleware in the list is the outermost.
7
8use std::fmt;
9use std::sync::Arc;
10
11use ferrin_spec::BoxFuture;
12use ferrin_spec::DynImageModel;
13use ferrin_spec::ImageModel;
14use ferrin_spec::ModelId;
15use ferrin_spec::ProviderId;
16use ferrin_spec::error::ProviderError;
17use ferrin_spec::image_model::ImageOptions;
18use ferrin_spec::image_model::ImageResult;
19
20/// The image model being wrapped.
21#[derive(Clone, Copy)]
22pub struct ImageMiddlewareContext<'a> {
23    /// The wrapped (inner) model.
24    pub model: &'a dyn DynImageModel,
25}
26
27impl fmt::Debug for ImageMiddlewareContext<'_> {
28    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
29        f.debug_struct("ImageMiddlewareContext")
30            .field("provider", self.model.provider())
31            .field("model_id", self.model.model_id())
32            .finish()
33    }
34}
35
36/// Continuation of a wrapped image `do_generate`.
37pub type ImageGenerateNext<'a> =
38    Box<dyn FnOnce(ImageOptions) -> BoxFuture<'a, Result<ImageResult, ProviderError>> + Send + 'a>;
39
40/// Intercepts image model calls. Every method has a pass-through default.
41///
42/// [`max_images_per_call`](Self::max_images_per_call) receives the wrapped
43/// model and returns the value the wrapper reports; its default forwards the
44/// inner model's value.
45pub trait ImageModelMiddleware: Send + Sync + 'static {
46    /// Rewrites the image options before the call.
47    fn transform_params<'a>(
48        &'a self,
49        options: ImageOptions,
50        _ctx: ImageMiddlewareContext<'a>,
51    ) -> BoxFuture<'a, Result<ImageOptions, ProviderError>> {
52        Box::pin(async move { Ok(options) })
53    }
54
55    /// Wraps `do_generate`.
56    fn wrap_generate<'a>(
57        &'a self,
58        options: ImageOptions,
59        next: ImageGenerateNext<'a>,
60        _ctx: ImageMiddlewareContext<'a>,
61    ) -> BoxFuture<'a, Result<ImageResult, ProviderError>> {
62        next(options)
63    }
64
65    /// Overrides the provider id reported by the wrapped model.
66    fn override_provider(&self, _model: &dyn DynImageModel) -> Option<ProviderId> {
67        None
68    }
69
70    /// Overrides the model id reported by the wrapped model.
71    fn override_model_id(&self, _model: &dyn DynImageModel) -> Option<ModelId> {
72        None
73    }
74
75    /// The per-call image limit reported by the wrapper; defaults to the
76    /// inner value.
77    fn max_images_per_call(&self, model: &dyn DynImageModel) -> Option<usize> {
78        model.max_images_per_call()
79    }
80}
81
82/// Wraps `model` with `middleware`; the first entry becomes the outermost
83/// layer. An empty list returns `model` unchanged.
84#[must_use]
85pub fn wrap_image_model(
86    model: Arc<dyn DynImageModel>,
87    middleware: impl IntoIterator<Item = Arc<dyn ImageModelMiddleware>, IntoIter: DoubleEndedIterator>,
88) -> Arc<dyn DynImageModel> {
89    middleware.into_iter().rev().fold(model, |inner, layer| {
90        let provider = layer
91            .override_provider(inner.as_ref())
92            .unwrap_or_else(|| inner.provider().clone());
93        let model_id = layer
94            .override_model_id(inner.as_ref())
95            .unwrap_or_else(|| inner.model_id().clone());
96        Arc::new(WrappedImageModel {
97            inner,
98            layer,
99            provider,
100            model_id,
101        })
102    })
103}
104
105struct WrappedImageModel {
106    inner: Arc<dyn DynImageModel>,
107    layer: Arc<dyn ImageModelMiddleware>,
108    provider: ProviderId,
109    model_id: ModelId,
110}
111
112impl fmt::Debug for WrappedImageModel {
113    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
114        f.debug_struct("WrappedImageModel")
115            .field("provider", &self.provider)
116            .field("model_id", &self.model_id)
117            .finish_non_exhaustive()
118    }
119}
120
121impl ImageModel for WrappedImageModel {
122    fn provider(&self) -> &ProviderId {
123        &self.provider
124    }
125
126    fn model_id(&self) -> &ModelId {
127        &self.model_id
128    }
129
130    fn max_images_per_call(&self) -> Option<usize> {
131        self.layer.max_images_per_call(self.inner.as_ref())
132    }
133
134    async fn do_generate(&self, options: ImageOptions) -> Result<ImageResult, ProviderError> {
135        let ctx = ImageMiddlewareContext {
136            model: self.inner.as_ref(),
137        };
138        let options = self.layer.transform_params(options, ctx).await?;
139        let inner = &self.inner;
140        self.layer
141            .wrap_generate(
142                options,
143                Box::new(move |options| inner.do_generate(options)),
144                ctx,
145            )
146            .await
147    }
148}