Skip to main content

ferrin_core/
image.rs

1//! Image generation: [`generate_image`] splits `n` images into provider
2//! calls, runs them concurrently and retries empty results.
3//!
4//! Design: `docs/01-architecture/11-other-modalities.md` ยง2.
5
6use std::future::IntoFuture;
7use std::sync::Arc;
8use std::sync::Mutex;
9
10use bytes::Bytes;
11use ferrin_provider_util::media_type::detect_media_type_for;
12use ferrin_spec::AspectRatio;
13use ferrin_spec::BoxFuture;
14use ferrin_spec::DynImageModel;
15use ferrin_spec::ImageModelRef;
16use ferrin_spec::ImageSize;
17use ferrin_spec::JsonValue;
18use ferrin_spec::MediaType;
19use ferrin_spec::ProviderMetadata;
20use ferrin_spec::ResponseMetadata;
21use ferrin_spec::Warning;
22use ferrin_spec::error::NoContentGeneratedError;
23use ferrin_spec::error::ProviderError;
24pub use ferrin_spec::image_model::ImageFile;
25use ferrin_spec::image_model::ImageOptions;
26use ferrin_spec::image_model::ImageResult;
27pub use ferrin_spec::image_model::ImageUsage;
28use tokio::task::JoinSet;
29use tokio_util::sync::CancellationToken;
30use tracing::Instrument;
31
32use crate::error::Error;
33use crate::modality::ModalityOptions;
34use crate::modality::add_optional;
35use crate::modality::impl_modality_builder;
36use crate::modality_metadata::merge_image_metadata;
37use crate::registry::ProviderRegistry;
38use crate::registry::default::resolve_model;
39use crate::retry::RetryPolicy;
40use crate::retry::retry_with;
41use crate::telemetry::ModelIdentity;
42use crate::telemetry::spans;
43
44/// Marker message of the internal retryable "empty result" error.
45const NO_IMAGE_MESSAGE: &str = "no image generated";
46
47/// Media type used when neither the provider nor detection knows better.
48const DEFAULT_IMAGE_MEDIA_TYPE: &str = "image/png";
49
50/// A generated image.
51#[derive(Debug, Clone, PartialEq, Eq)]
52pub struct GeneratedImage {
53    /// Image bytes.
54    pub data: Bytes,
55    /// Media type (reported, detected, or `image/png`).
56    pub media_type: MediaType,
57    /// Provider metadata of this image (the entry of the provider's
58    /// `images` list at this index).
59    pub provider_metadata: Option<ProviderMetadata>,
60}
61
62impl GeneratedImage {
63    /// The image bytes as base64.
64    #[must_use]
65    pub fn base64(&self) -> String {
66        use base64::Engine as _;
67        base64::engine::general_purpose::STANDARD.encode(&self.data)
68    }
69}
70
71/// One provider call of [`generate_image`].
72#[derive(Debug, Clone, PartialEq)]
73pub struct ImageCall {
74    /// Images of this call.
75    pub images: Vec<GeneratedImage>,
76    /// Warnings of this call.
77    pub warnings: Vec<Warning>,
78    /// Response metadata of this call.
79    pub response: ResponseMetadata,
80    /// Provider metadata of this call.
81    pub provider_metadata: Option<ProviderMetadata>,
82    /// Usage of this call.
83    pub usage: Option<ImageUsage>,
84}
85
86/// Result of [`generate_image`].
87#[derive(Debug, Clone, PartialEq)]
88pub struct GenerateImageResult {
89    /// All images, in call order.
90    pub images: Vec<GeneratedImage>,
91    /// The individual provider calls.
92    pub calls: Vec<ImageCall>,
93    /// Warnings of all calls.
94    pub warnings: Vec<Warning>,
95    /// Response metadata of all calls.
96    pub responses: Vec<ResponseMetadata>,
97    /// Provider metadata merged over all calls.
98    pub provider_metadata: ProviderMetadata,
99    /// Usage summed over all calls.
100    pub usage: ImageUsage,
101}
102
103impl GenerateImageResult {
104    /// The first image.
105    #[must_use]
106    pub fn image(&self) -> Option<&GeneratedImage> {
107        self.images.first()
108    }
109}
110
111/// Generates images from a text prompt.
112#[must_use]
113pub fn generate_image(model: impl Into<ImageModelRef>, prompt: impl Into<String>) -> GenerateImage {
114    GenerateImage {
115        model: model.into(),
116        prompt: Some(prompt.into()),
117        n: 1,
118        max_images_per_call: None,
119        size: None,
120        aspect_ratio: None,
121        seed: None,
122        files: Vec::new(),
123        mask: None,
124        base: ModalityOptions::default(),
125    }
126}
127
128/// Generates images from reference images (image-to-image), optionally
129/// guided by a prompt set with [`GenerateImage::prompt`].
130#[must_use]
131pub fn edit_image(model: impl Into<ImageModelRef>, files: Vec<ImageFile>) -> GenerateImage {
132    GenerateImage {
133        model: model.into(),
134        prompt: None,
135        n: 1,
136        max_images_per_call: None,
137        size: None,
138        aspect_ratio: None,
139        seed: None,
140        files,
141        mask: None,
142        base: ModalityOptions::default(),
143    }
144}
145
146/// Builder returned by [`generate_image`]; `.await` runs the calls.
147#[derive(Debug)]
148pub struct GenerateImage {
149    model: ImageModelRef,
150    prompt: Option<String>,
151    n: u32,
152    max_images_per_call: Option<u32>,
153    size: Option<ImageSize>,
154    aspect_ratio: Option<AspectRatio>,
155    seed: Option<u64>,
156    files: Vec<ImageFile>,
157    mask: Option<ImageFile>,
158    base: ModalityOptions,
159}
160
161impl GenerateImage {
162    /// Sets the text prompt.
163    #[must_use]
164    pub fn prompt(mut self, prompt: impl Into<String>) -> Self {
165        self.prompt = Some(prompt.into());
166        self
167    }
168
169    /// Number of images to generate (default 1).
170    #[must_use]
171    pub fn n(mut self, n: u32) -> Self {
172        self.n = n;
173        self
174    }
175
176    /// Overrides the model's images-per-call limit.
177    #[must_use]
178    pub fn max_images_per_call(mut self, max_images_per_call: u32) -> Self {
179        self.max_images_per_call = Some(max_images_per_call);
180        self
181    }
182
183    /// Image size (`1024x1024`).
184    #[must_use]
185    pub fn size(mut self, size: ImageSize) -> Self {
186        self.size = Some(size);
187        self
188    }
189
190    /// Aspect ratio (`16:9`).
191    #[must_use]
192    pub fn aspect_ratio(mut self, aspect_ratio: AspectRatio) -> Self {
193        self.aspect_ratio = Some(aspect_ratio);
194        self
195    }
196
197    /// Seed for reproducible generation.
198    #[must_use]
199    pub fn seed(mut self, seed: u64) -> Self {
200        self.seed = Some(seed);
201        self
202    }
203
204    /// Reference images.
205    #[must_use]
206    pub fn files(mut self, files: Vec<ImageFile>) -> Self {
207        self.files = files;
208        self
209    }
210
211    /// Adds one reference image.
212    #[must_use]
213    pub fn file(mut self, file: ImageFile) -> Self {
214        self.files.push(file);
215        self
216    }
217
218    /// Mask for inpainting.
219    #[must_use]
220    pub fn mask(mut self, mask: ImageFile) -> Self {
221        self.mask = Some(mask);
222        self
223    }
224}
225
226impl_modality_builder!(GenerateImage);
227
228impl IntoFuture for GenerateImage {
229    type Output = Result<GenerateImageResult, Error>;
230    type IntoFuture = BoxFuture<'static, Self::Output>;
231
232    fn into_future(self) -> Self::IntoFuture {
233        Box::pin(run(self))
234    }
235}
236
237/// Provider metadata of one image: the `images[index]` entry of every
238/// provider that lists per-image metadata.
239pub(crate) fn image_provider_metadata(
240    provider_metadata: Option<&ProviderMetadata>,
241    index: usize,
242) -> Option<ProviderMetadata> {
243    let mut result: Option<ProviderMetadata> = None;
244    for (provider, metadata) in provider_metadata? {
245        if let Some(JsonValue::Array(images)) = metadata.get("images")
246            && let Some(JsonValue::Object(entry)) = images.get(index)
247        {
248            result
249                .get_or_insert_with(ProviderMetadata::new)
250                .insert(provider.clone(), entry.clone());
251        }
252    }
253    result
254}
255
256/// Converts the images of one provider result.
257pub(crate) fn convert_images(result: &ImageResult) -> Vec<GeneratedImage> {
258    result
259        .images
260        .iter()
261        .enumerate()
262        .map(|(index, image)| GeneratedImage {
263            data: image.data.clone(),
264            media_type: image
265                .media_type
266                .clone()
267                .or_else(|| detect_media_type_for(&image.data, "image"))
268                .unwrap_or_else(|| MediaType::new(DEFAULT_IMAGE_MEDIA_TYPE)),
269            provider_metadata: image_provider_metadata(result.provider_metadata.as_ref(), index),
270        })
271        .collect()
272}
273
274/// Adds two optional usages field by field.
275pub(crate) fn add_image_usage(total: ImageUsage, usage: &ImageUsage) -> ImageUsage {
276    ImageUsage {
277        input_tokens: add_optional(total.input_tokens, usage.input_tokens),
278        output_tokens: add_optional(total.output_tokens, usage.output_tokens),
279        total_tokens: add_optional(total.total_tokens, usage.total_tokens),
280    }
281}
282
283fn no_image_error() -> ProviderError {
284    ProviderError::NoContentGenerated(NoContentGeneratedError::with_message(NO_IMAGE_MESSAGE))
285}
286
287fn is_no_image(error: &ProviderError) -> bool {
288    matches!(error, ProviderError::NoContentGenerated(inner) if inner.message == NO_IMAGE_MESSAGE)
289}
290
291fn ended_without_image(error: &Error) -> bool {
292    match error {
293        Error::Provider(error) => is_no_image(error),
294        Error::Retry { errors, .. } => errors.last().is_some_and(is_no_image),
295        _ => false,
296    }
297}
298
299/// One provider call (owned so that it can run on a task).
300struct ImageCallTask {
301    model: Arc<dyn DynImageModel>,
302    options: ImageOptions,
303    retry_policy: RetryPolicy,
304    cancellation: CancellationToken,
305}
306
307impl ImageCallTask {
308    /// Runs the call with retries; returns every provider result received
309    /// (empty results are retried and kept for their metadata).
310    async fn run(self) -> Result<Vec<ImageResult>, Error> {
311        let results: Arc<Mutex<Vec<ImageResult>>> = Arc::new(Mutex::new(Vec::new()));
312        let outcome = retry_with(
313            &self.retry_policy,
314            &self.cancellation,
315            |error| error.is_retryable() || is_no_image(error),
316            |_| {
317                let results = Arc::clone(&results);
318                let mut options = self.options.clone();
319                options.cancellation = self.cancellation.child_token();
320                let model = Arc::clone(&self.model);
321                async move {
322                    let result = model.do_generate(options).await.map_err(Error::from)?;
323                    let empty = result.images.is_empty() && result.is_retryable != Some(false);
324                    results
325                        .lock()
326                        .unwrap_or_else(std::sync::PoisonError::into_inner)
327                        .push(result);
328                    if empty {
329                        Err(Error::from(no_image_error()))
330                    } else {
331                        Ok(())
332                    }
333                }
334            },
335        )
336        .await;
337        match outcome {
338            Ok(()) => {}
339            Err(error) if ended_without_image(&error) => {}
340            Err(error) => return Err(error),
341        }
342        let collected = std::mem::take(
343            &mut *results
344                .lock()
345                .unwrap_or_else(std::sync::PoisonError::into_inner),
346        );
347        Ok(collected)
348    }
349}
350
351async fn run(builder: GenerateImage) -> Result<GenerateImageResult, Error> {
352    if builder.n == 0 {
353        return Err(Error::invalid_argument("n", "must be at least 1"));
354    }
355    if builder.max_images_per_call == Some(0) {
356        return Err(Error::invalid_argument(
357            "max_images_per_call",
358            "must be at least 1",
359        ));
360    }
361    let model = resolve_model(&builder.model, ProviderRegistry::image_model)?;
362    let identity = ModelIdentity::new(model.provider().clone(), model.model_id().clone());
363    let span = spans::modality_span("image", &identity);
364    let base = builder.base.clone();
365    base.run(|_, token| run_calls(model, identity, builder, token).instrument(span))
366        .await
367}
368
369async fn run_calls(
370    model: Arc<dyn DynImageModel>,
371    identity: ModelIdentity,
372    builder: GenerateImage,
373    cancellation: CancellationToken,
374) -> Result<GenerateImageResult, Error> {
375    let per_call = builder
376        .max_images_per_call
377        .map(|limit| usize::try_from(limit).unwrap_or(usize::MAX))
378        .or_else(|| model.max_images_per_call().filter(|limit| *limit > 0))
379        .unwrap_or(1);
380    let n = usize::try_from(builder.n).unwrap_or(usize::MAX);
381    let call_count = n.div_ceil(per_call);
382    let counts: Vec<u32> = (0..call_count)
383        .map(|index| {
384            let remaining = n.saturating_sub(index.saturating_mul(per_call));
385            u32::try_from(remaining.min(per_call)).unwrap_or(u32::MAX)
386        })
387        .collect();
388
389    let template = ImageOptions {
390        prompt: builder.prompt.clone(),
391        n: 1,
392        size: builder.size,
393        aspect_ratio: builder.aspect_ratio,
394        seed: builder.seed,
395        files: builder.files.clone(),
396        mask: builder.mask.clone(),
397        provider_options: builder.base.provider_options.clone(),
398        headers: builder.base.request_headers(),
399        cancellation: cancellation.clone(),
400    };
401    let make_task = |count: u32| ImageCallTask {
402        model: Arc::clone(&model),
403        options: ImageOptions {
404            n: count,
405            ..template.clone()
406        },
407        retry_policy: builder.base.retry_policy.clone(),
408        cancellation: cancellation.clone(),
409    };
410
411    let mut groups: Vec<Option<Vec<ImageResult>>> = (0..counts.len()).map(|_| None).collect();
412    if let [count] = counts.as_slice() {
413        groups = vec![Some(make_task(*count).run().await?)];
414    } else {
415        let mut tasks: JoinSet<(usize, Result<Vec<ImageResult>, Error>)> = JoinSet::new();
416        for (index, count) in counts.iter().enumerate() {
417            let task = make_task(*count);
418            tasks.spawn(async move { (index, task.run().await) });
419        }
420        while let Some(joined) = tasks.join_next().await {
421            let (index, result) =
422                joined.map_err(|error| Error::message(format!("image task failed: {error}")))?;
423            if let Some(slot) = groups.get_mut(index) {
424                *slot = Some(result?);
425            }
426        }
427    }
428
429    let mut images: Vec<GeneratedImage> = Vec::new();
430    let mut calls: Vec<ImageCall> = Vec::new();
431    let mut warnings: Vec<Warning> = Vec::new();
432    let mut responses: Vec<ResponseMetadata> = Vec::new();
433    let mut provider_metadata = ProviderMetadata::new();
434    let mut usage = ImageUsage::default();
435    for result in groups.into_iter().flatten().flatten() {
436        let call_images = convert_images(&result);
437        images.extend(call_images.iter().cloned());
438        warnings.extend(result.warnings.iter().cloned());
439        responses.push(result.response.clone());
440        if let Some(call_usage) = &result.usage {
441            usage = add_image_usage(usage, call_usage);
442        }
443        if let Some(metadata) = &result.provider_metadata {
444            merge_image_metadata(&mut provider_metadata, metadata);
445        }
446        calls.push(ImageCall {
447            images: call_images,
448            warnings: result.warnings,
449            response: result.response,
450            provider_metadata: result.provider_metadata,
451            usage: result.usage,
452        });
453    }
454    if images.is_empty() {
455        return Err(Error::NoImageGenerated { responses });
456    }
457    spans::log_warnings(&warnings, &identity);
458    Ok(GenerateImageResult {
459        images,
460        calls,
461        warnings,
462        responses,
463        provider_metadata,
464        usage,
465    })
466}