ferrin_core/middleware/
image.rs1use 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#[derive(Clone, Copy)]
22pub struct ImageMiddlewareContext<'a> {
23 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
36pub type ImageGenerateNext<'a> =
38 Box<dyn FnOnce(ImageOptions) -> BoxFuture<'a, Result<ImageResult, ProviderError>> + Send + 'a>;
39
40pub trait ImageModelMiddleware: Send + Sync + 'static {
46 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 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 fn override_provider(&self, _model: &dyn DynImageModel) -> Option<ProviderId> {
67 None
68 }
69
70 fn override_model_id(&self, _model: &dyn DynImageModel) -> Option<ModelId> {
72 None
73 }
74
75 fn max_images_per_call(&self, model: &dyn DynImageModel) -> Option<usize> {
78 model.max_images_per_call()
79 }
80}
81
82#[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}