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;