Skip to main content

rig_core/
image_generation.rs

1//! Everything related to core image generation abstractions in Rig.
2//! Rig allows calling a number of different providers (that support image generation) using the [ImageGenerationModel] trait.
3use crate::markers::{Missing, Provided};
4use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
5use serde_json::Value;
6
7crate::provider_response::provider_error_enum!(
8    ImageGenerationError, "image generation" {
9        /// Error building the image generation request
10        #[error("RequestError: {0}")]
11        RequestError(#[from] Box<dyn std::error::Error + Send + Sync + 'static>),
12    }
13);
14
15/// A unified response for a model image generation, returning both the image and the raw response.
16#[derive(Debug)]
17pub struct ImageGenerationResponse<T> {
18    pub image: Vec<u8>,
19    pub response: T,
20}
21
22pub trait ImageGenerationModel: Clone + WasmCompatSend + WasmCompatSync {
23    type Response: WasmCompatSend + WasmCompatSync;
24
25    type Client;
26
27    fn make(client: &Self::Client, model: impl Into<String>) -> Self;
28
29    fn image_generation(
30        &self,
31        request: ImageGenerationRequest,
32    ) -> impl std::future::Future<
33        Output = Result<ImageGenerationResponse<Self::Response>, ImageGenerationError>,
34    > + WasmCompatSend;
35
36    fn image_generation_request(&self) -> ImageGenerationRequestBuilder<Self, Missing> {
37        ImageGenerationRequestBuilder::new(self.clone())
38    }
39}
40/// An image generation request.
41pub struct ImageGenerationRequest {
42    pub prompt: String,
43    pub width: u32,
44    pub height: u32,
45    pub additional_params: Option<Value>,
46}
47
48/// A builder for `ImageGenerationRequest`.
49/// Can be sent to a model provider.
50pub struct ImageGenerationRequestBuilder<M, P = Missing>
51where
52    M: ImageGenerationModel,
53{
54    model: M,
55    prompt: P,
56    width: u32,
57    height: u32,
58    additional_params: Option<Value>,
59}
60
61impl<M> ImageGenerationRequestBuilder<M, Missing>
62where
63    M: ImageGenerationModel,
64{
65    pub fn new(model: M) -> Self {
66        Self {
67            model,
68            prompt: Missing,
69            height: 256,
70            width: 256,
71            additional_params: None,
72        }
73    }
74}
75
76impl<M, P> ImageGenerationRequestBuilder<M, P>
77where
78    M: ImageGenerationModel,
79{
80    /// Sets the prompt for the image generation request
81    pub fn prompt(self, prompt: &str) -> ImageGenerationRequestBuilder<M, Provided<String>> {
82        ImageGenerationRequestBuilder {
83            model: self.model,
84            prompt: Provided(prompt.to_string()),
85            width: self.width,
86            height: self.height,
87            additional_params: self.additional_params,
88        }
89    }
90
91    /// The width of the generated image
92    pub fn width(mut self, width: u32) -> Self {
93        self.width = width;
94        self
95    }
96
97    /// The height of the generated image
98    pub fn height(mut self, height: u32) -> Self {
99        self.height = height;
100        self
101    }
102
103    /// Adds additional parameters to the image generation request.
104    pub fn additional_params(mut self, params: Value) -> Self {
105        self.additional_params = Some(params);
106        self
107    }
108}
109
110impl<M> ImageGenerationRequestBuilder<M, Provided<String>>
111where
112    M: ImageGenerationModel,
113{
114    pub fn build(self) -> ImageGenerationRequest {
115        ImageGenerationRequest {
116            prompt: self.prompt.0,
117            width: self.width,
118            height: self.height,
119            additional_params: self.additional_params,
120        }
121    }
122
123    pub async fn send(self) -> Result<ImageGenerationResponse<M::Response>, ImageGenerationError> {
124        let model = self.model.clone();
125
126        model.image_generation(self.build()).await
127    }
128}
129
130#[cfg(test)]
131mod provider_response_tests {
132    use super::*;
133    use crate::{http_client, provider_response};
134    use http::StatusCode;
135
136    #[test]
137    fn image_generation_error_provider_response_helpers_with_preserved_json_body() {
138        let body = r#"{"error":{"message":"content policy"}}"#;
139        let error = ImageGenerationError::ProviderResponse(
140            provider_response::ProviderResponseError::without_status(body.to_string()),
141        );
142
143        assert_eq!(error.provider_response_body(), Some(body));
144        assert_eq!(error.provider_response_status(), None);
145        assert_eq!(
146            error.provider_response_json().expect("valid JSON"),
147            Some(serde_json::json!({ "error": { "message": "content policy" } }))
148        );
149    }
150
151    #[test]
152    fn image_generation_error_provider_response_helpers_with_http_non_success() {
153        let body = r#"{"error":{"message":"bad request"}}"#;
154        let error =
155            ImageGenerationError::HttpError(http_client::Error::InvalidStatusCodeWithMessage(
156                StatusCode::BAD_REQUEST,
157                body.to_string(),
158            ));
159
160        assert_eq!(error.provider_response_body(), Some(body));
161        assert_eq!(
162            error.provider_response_status(),
163            Some(StatusCode::BAD_REQUEST)
164        );
165        assert_eq!(
166            error.provider_response_json().expect("valid JSON"),
167            Some(serde_json::json!({ "error": { "message": "bad request" } }))
168        );
169    }
170
171    #[test]
172    fn image_generation_error_provider_error_is_not_a_provider_response() {
173        let error = ImageGenerationError::ProviderError("internal diagnostic".to_string());
174
175        assert_eq!(error.provider_response_body(), None);
176        assert_eq!(error.provider_response_status(), None);
177        assert_eq!(error.provider_response_json().expect("no body"), None);
178    }
179
180    #[test]
181    fn image_generation_error_provider_response_helpers_with_unrelated_variant() {
182        let error = ImageGenerationError::ResponseError("parse failed".to_string());
183
184        assert_eq!(error.provider_response_body(), None);
185        assert_eq!(error.provider_response_status(), None);
186        assert_eq!(error.provider_response_json().expect("no body"), None);
187    }
188}