1use 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::merge_provider_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
44const NO_IMAGE_MESSAGE: &str = "no image generated";
46
47const DEFAULT_IMAGE_MEDIA_TYPE: &str = "image/png";
49
50#[derive(Debug, Clone, PartialEq, Eq)]
52pub struct GeneratedImage {
53 pub data: Bytes,
55 pub media_type: MediaType,
57 pub provider_metadata: Option<ProviderMetadata>,
60}
61
62impl GeneratedImage {
63 #[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#[derive(Debug, Clone, PartialEq)]
73pub struct ImageCall {
74 pub images: Vec<GeneratedImage>,
76 pub warnings: Vec<Warning>,
78 pub response: ResponseMetadata,
80 pub provider_metadata: Option<ProviderMetadata>,
82 pub usage: Option<ImageUsage>,
84}
85
86#[derive(Debug, Clone, PartialEq)]
88pub struct GenerateImageResult {
89 pub images: Vec<GeneratedImage>,
91 pub calls: Vec<ImageCall>,
93 pub warnings: Vec<Warning>,
95 pub responses: Vec<ResponseMetadata>,
97 pub provider_metadata: ProviderMetadata,
99 pub usage: ImageUsage,
101}
102
103impl GenerateImageResult {
104 #[must_use]
106 pub fn image(&self) -> Option<&GeneratedImage> {
107 self.images.first()
108 }
109}
110
111#[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#[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#[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 #[must_use]
164 pub fn prompt(mut self, prompt: impl Into<String>) -> Self {
165 self.prompt = Some(prompt.into());
166 self
167 }
168
169 #[must_use]
171 pub fn n(mut self, n: u32) -> Self {
172 self.n = n;
173 self
174 }
175
176 #[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 #[must_use]
185 pub fn size(mut self, size: ImageSize) -> Self {
186 self.size = Some(size);
187 self
188 }
189
190 #[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 #[must_use]
199 pub fn seed(mut self, seed: u64) -> Self {
200 self.seed = Some(seed);
201 self
202 }
203
204 #[must_use]
206 pub fn files(mut self, files: Vec<ImageFile>) -> Self {
207 self.files = files;
208 self
209 }
210
211 #[must_use]
213 pub fn file(mut self, file: ImageFile) -> Self {
214 self.files.push(file);
215 self
216 }
217
218 #[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
237pub(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
256pub(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
274pub(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
299struct ImageCallTask {
301 model: Arc<dyn DynImageModel>,
302 options: ImageOptions,
303 retry_policy: RetryPolicy,
304 cancellation: CancellationToken,
305}
306
307impl ImageCallTask {
308 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_provider_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}