Skip to main content

rig_core/
image_generation.rs

1//! Image-generation requests, normalized responses, and model interfaces.
2//!
3//! ```no_run
4//! use rig_core::DynModel;
5//! use rig_core::image_generation::ImageGenerationRequestBuilder;
6//! use rig_core::operation::ImageGeneration;
7//!
8//! # async fn example(model: DynModel<ImageGeneration>) -> Result<(), Box<dyn std::error::Error>> {
9//! let request = ImageGenerationRequestBuilder::new("A mountain lake")
10//!     .width(1024)
11//!     .height(1024)
12//!     .build();
13//! let response = model.call(request).await?;
14//! # let _ = response;
15//! # Ok(())
16//! # }
17//! ```
18use crate::completion::Usage;
19use crate::error::ProviderError;
20use serde::{Deserialize, Serialize};
21use serde_json::Value;
22
23/// Generated image bytes and normalized provider metadata.
24#[derive(Debug, Clone, Serialize, Deserialize)]
25pub struct ImageGenerationResponse {
26    /// The generated image, decoded to bytes.
27    pub image: Vec<u8>,
28    /// Usage as the provider reported it; every counter is `None` when the
29    /// provider reported none (see [`Usage`]).
30    #[serde(default)]
31    pub usage: Usage,
32    /// Stable descriptor name of the provider that produced this response,
33    /// for example `"openai"`. Always populated.
34    pub provider: String,
35    /// Provider-reported model identifier, when the wire response named one.
36    #[serde(default)]
37    pub model: Option<String>,
38    /// Provider-assigned response-scoped identifier, when reported.
39    #[serde(default, skip_serializing_if = "Option::is_none")]
40    pub response_id: Option<String>,
41    /// Transport request ID from HTTP headers, or `None` when unreported.
42    #[serde(default, skip_serializing_if = "Option::is_none")]
43    pub provider_request_id: Option<String>,
44    /// Provider response metadata. May be null for byte-only responses or
45    /// responses constructed without metadata; image bytes remain in [`Self::image`].
46    #[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
47    pub raw: serde_json::Value,
48}
49
50impl ImageGenerationResponse {
51    /// A response carrying `image`. The driver writes the provider, the
52    /// transport request id and the reply document; decoders set what the
53    /// provider reported.
54    pub fn new(image: Vec<u8>) -> Self {
55        Self {
56            image,
57            usage: Usage::default(),
58            provider: String::new(),
59            model: None,
60            response_id: None,
61            provider_request_id: None,
62            raw: serde_json::Value::Null,
63        }
64    }
65}
66
67/// Normalizes provider image payloads. The driver writes the provider, request
68/// id and `raw` afterwards.
69pub trait NormalizeImageGenerationResponse {
70    /// Normalize this payload.
71    fn normalize(self) -> Result<ImageGenerationResponse, ProviderError>;
72}
73
74pub struct ImageGenerationRequest {
75    pub prompt: String,
76    pub width: u32,
77    pub height: u32,
78    pub additional_params: Option<Value>,
79}
80
81/// Builds an image request for a prompt. Defaults to 256 by 256 pixels;
82/// supported dimensions depend on the provider.
83pub struct ImageGenerationRequestBuilder {
84    request: ImageGenerationRequest,
85}
86
87impl ImageGenerationRequestBuilder {
88    /// A request for `prompt`.
89    pub fn new(prompt: impl Into<String>) -> Self {
90        Self {
91            request: ImageGenerationRequest {
92                prompt: prompt.into(),
93                width: 256,
94                height: 256,
95                additional_params: None,
96            },
97        }
98    }
99
100    /// The width of the generated image.
101    pub fn width(mut self, width: u32) -> Self {
102        self.request.width = width;
103        self
104    }
105
106    /// The height of the generated image.
107    pub fn height(mut self, height: u32) -> Self {
108        self.request.height = height;
109        self
110    }
111
112    /// Merges provider-specific parameters over earlier ones, key by key for
113    /// JSON objects; `None` clears existing parameters.
114    pub fn additional_params(mut self, params: impl Into<Option<Value>>) -> Self {
115        self.request.additional_params =
116            crate::json_utils::merge_params(self.request.additional_params.take(), params.into());
117        self
118    }
119
120    /// Builds the image generation request.
121    pub fn build(self) -> ImageGenerationRequest {
122        self.request
123    }
124}
125
126#[cfg(test)]
127mod builder_tests;
128#[cfg(test)]
129mod provider_response_tests;