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