Skip to main content

rig_core/providers/gemini/
image_generation.rs

1//! Gemini image generation support.
2
3use super::client::{ApiResponse, Client};
4use super::completion::gemini_api_types::{
5    Content, GenerateContentRequest, GenerateContentResponse, GenerationConfig, ImageConfig, Part,
6    PartKind, ResponseModality, Role,
7};
8use crate::http_client::HttpClientExt;
9use crate::image_generation::{ImageGenerationError, ImageGenerationRequest};
10use crate::{http_client, image_generation};
11use base64::Engine;
12use base64::prelude::BASE64_STANDARD;
13use serde_json::Value;
14
15/// `gemini-2.5-flash-image` image generation model, commonly referred to as Nano Banana.
16pub const GEMINI_2_5_FLASH_IMAGE: &str = super::completion::GEMINI_2_5_FLASH_IMAGE;
17
18/// Gemini image generation model.
19#[derive(Clone)]
20pub struct ImageGenerationModel<T = reqwest::Client> {
21    client: Client<T>,
22    /// Name of the model, for example [`GEMINI_2_5_FLASH_IMAGE`].
23    pub model: String,
24}
25
26impl<T> ImageGenerationModel<T> {
27    pub(crate) fn new(client: Client<T>, model: impl Into<String>) -> Self {
28        Self {
29            client,
30            model: model.into(),
31        }
32    }
33}
34
35impl TryFrom<GenerateContentResponse>
36    for image_generation::ImageGenerationResponse<GenerateContentResponse>
37{
38    type Error = ImageGenerationError;
39
40    fn try_from(value: GenerateContentResponse) -> Result<Self, Self::Error> {
41        let image = first_image_bytes(&value)?;
42
43        Ok(image_generation::ImageGenerationResponse {
44            image,
45            response: value,
46        })
47    }
48}
49
50impl<T> image_generation::ImageGenerationModel for ImageGenerationModel<T>
51where
52    T: HttpClientExt + Clone + Default + std::fmt::Debug + Send + 'static,
53{
54    type Response = GenerateContentResponse;
55
56    type Client = Client<T>;
57
58    fn make(client: &Self::Client, model: impl Into<String>) -> Self {
59        Self::new(client.clone(), model)
60    }
61
62    async fn image_generation(
63        &self,
64        generation_request: ImageGenerationRequest,
65    ) -> Result<image_generation::ImageGenerationResponse<Self::Response>, ImageGenerationError>
66    {
67        let body = serde_json::to_vec(&create_request_body(generation_request)?)?;
68
69        let request = self
70            .client
71            .post(generate_content_path(&self.model))?
72            .body(body)
73            .map_err(|e| ImageGenerationError::HttpError(e.into()))?;
74
75        let response = self.client.send(request).await?;
76
77        let status = response.status();
78        let text = http_client::text(response).await?;
79
80        if !status.is_success() {
81            return Err(ImageGenerationError::from_http_response(status, text));
82        }
83
84        match serde_json::from_str::<ApiResponse<GenerateContentResponse>>(&text)? {
85            ApiResponse::Ok(response) => response.try_into(),
86            ApiResponse::Err(err) => {
87                tracing::warn!(message = %err.error.message, "provider returned an error response");
88                Err(ImageGenerationError::from_http_response(status, text))
89            }
90        }
91    }
92}
93
94fn generate_content_path(model: &str) -> String {
95    format!("/v1beta/models/{model}:generateContent")
96}
97
98fn create_request_body(
99    generation_request: ImageGenerationRequest,
100) -> Result<Value, ImageGenerationError> {
101    let request = GenerateContentRequest {
102        contents: vec![Content {
103            role: Some(Role::User),
104            parts: vec![Part {
105                thought: None,
106                thought_signature: None,
107                part: PartKind::Text(generation_request.prompt),
108                additional_params: None,
109            }],
110        }],
111        tools: None,
112        tool_config: None,
113        generation_config: Some(GenerationConfig {
114            response_modalities: Some(vec![ResponseModality::Image]),
115            image_config: Some(ImageConfig {
116                aspect_ratio: aspect_ratio(generation_request.width, generation_request.height),
117                image_size: None,
118            }),
119            ..Default::default()
120        }),
121        safety_settings: None,
122        system_instruction: None,
123        additional_params: None,
124    };
125
126    let mut body = serde_json::to_value(request)?;
127
128    if let Some(additional_params) = generation_request.additional_params {
129        merge_json_deep(&mut body, additional_params);
130    }
131
132    Ok(body)
133}
134
135fn merge_json_deep(target: &mut Value, source: Value) {
136    match (target, source) {
137        (Value::Object(target), Value::Object(source)) => {
138            for (key, value) in source {
139                if let Some(existing) = target.get_mut(&key) {
140                    merge_json_deep(existing, value);
141                } else {
142                    target.insert(key, value);
143                }
144            }
145        }
146        (target, source) => *target = source,
147    }
148}
149
150fn aspect_ratio(width: u32, height: u32) -> Option<String> {
151    match (width, height) {
152        (0, _) | (_, 0) => None,
153        (w, h) if w == h => Some("1:1".to_string()),
154        (w, h) if w.saturating_mul(3) == h.saturating_mul(4) => Some("3:4".to_string()),
155        (w, h) if w.saturating_mul(4) == h.saturating_mul(3) => Some("4:3".to_string()),
156        (w, h) if w.saturating_mul(9) == h.saturating_mul(16) => Some("9:16".to_string()),
157        (w, h) if w.saturating_mul(16) == h.saturating_mul(9) => Some("16:9".to_string()),
158        _ => None,
159    }
160}
161
162fn first_image_bytes(response: &GenerateContentResponse) -> Result<Vec<u8>, ImageGenerationError> {
163    for candidate in &response.candidates {
164        let Some(content) = &candidate.content else {
165            continue;
166        };
167
168        for part in &content.parts {
169            if part.thought == Some(true) {
170                continue;
171            }
172
173            if let PartKind::InlineData(inline_data) = &part.part {
174                if !inline_data.mime_type.starts_with("image/") {
175                    continue;
176                }
177
178                return BASE64_STANDARD.decode(&inline_data.data).map_err(|err| {
179                    ImageGenerationError::ResponseError(format!(
180                        "Gemini image data was not valid base64: {err}"
181                    ))
182                });
183            }
184        }
185    }
186
187    Err(ImageGenerationError::ResponseError(
188        "Gemini image generation response did not include image data".into(),
189    ))
190}
191
192#[cfg(test)]
193mod tests {
194    use super::*;
195    use crate::providers::gemini::completion::gemini_api_types::{
196        Blob, ContentCandidate, FinishReason, UsageMetadata,
197    };
198    use serde_json::json;
199
200    fn image_generation_request(prompt: &str) -> ImageGenerationRequest {
201        ImageGenerationRequest {
202            prompt: prompt.to_string(),
203            width: 1024,
204            height: 1024,
205            additional_params: None,
206        }
207    }
208
209    #[test]
210    fn request_body_uses_gemini_image_generation_shape() {
211        let body = create_request_body(image_generation_request("Generate an image of an axolotl"))
212            .expect("request should serialize");
213
214        assert_eq!(
215            generate_content_path(GEMINI_2_5_FLASH_IMAGE),
216            "/v1beta/models/gemini-2.5-flash-image:generateContent"
217        );
218        assert_eq!(body["contents"][0]["role"], "user");
219        assert_eq!(
220            body["contents"][0]["parts"][0]["text"],
221            "Generate an image of an axolotl"
222        );
223        assert_eq!(
224            body["generationConfig"]["responseModalities"],
225            json!(["IMAGE"])
226        );
227        assert_eq!(
228            body["generationConfig"]["imageConfig"]["aspectRatio"],
229            "1:1"
230        );
231    }
232
233    #[test]
234    fn request_body_allows_additional_params_to_override_image_config() {
235        let mut request = image_generation_request("Generate an image of an axolotl");
236        request.additional_params = Some(json!({
237            "generationConfig": {
238                "imageConfig": {
239                    "aspectRatio": "16:9",
240                    "imageSize": "2K"
241                }
242            }
243        }));
244
245        let body = create_request_body(request).expect("request should serialize");
246
247        assert_eq!(
248            body["generationConfig"]["imageConfig"]["aspectRatio"],
249            "16:9"
250        );
251        assert_eq!(body["generationConfig"]["imageConfig"]["imageSize"], "2K");
252        assert_eq!(
253            body["generationConfig"]["responseModalities"],
254            json!(["IMAGE"])
255        );
256    }
257
258    #[test]
259    fn response_parsing_returns_first_non_thought_inline_image() {
260        let response = GenerateContentResponse {
261            candidates: vec![ContentCandidate {
262                content: Some(Content {
263                    role: Some(Role::Model),
264                    parts: vec![
265                        Part {
266                            thought: Some(false),
267                            thought_signature: None,
268                            part: PartKind::Text("Here you go".to_string()),
269                            additional_params: None,
270                        },
271                        Part {
272                            thought: Some(true),
273                            thought_signature: None,
274                            part: PartKind::InlineData(Blob {
275                                mime_type: "image/png".to_string(),
276                                data: BASE64_STANDARD.encode("thought image"),
277                            }),
278                            additional_params: None,
279                        },
280                        Part {
281                            thought: Some(false),
282                            thought_signature: None,
283                            part: PartKind::InlineData(Blob {
284                                mime_type: "image/png".to_string(),
285                                data: BASE64_STANDARD.encode("final image"),
286                            }),
287                            additional_params: None,
288                        },
289                    ],
290                }),
291                finish_reason: Some(FinishReason::Stop),
292                safety_ratings: None,
293                citation_metadata: None,
294                token_count: None,
295                avg_logprobs: None,
296                logprobs_result: None,
297                index: None,
298                finish_message: None,
299            }],
300            prompt_feedback: None,
301            usage_metadata: Some(UsageMetadata {
302                prompt_token_count: 1,
303                cached_content_token_count: None,
304                candidates_token_count: Some(1),
305                total_token_count: 2,
306                thoughts_token_count: None,
307                prompt_tokens_details: None,
308                cache_tokens_details: None,
309                candidates_tokens_details: None,
310                tool_use_prompt_token_count: None,
311                tool_use_prompt_tokens_details: None,
312                traffic_type: None,
313            }),
314            model_version: Some(GEMINI_2_5_FLASH_IMAGE.to_string()),
315            response_id: "response-id".to_string(),
316        };
317
318        let parsed: image_generation::ImageGenerationResponse<GenerateContentResponse> = response
319            .try_into()
320            .expect("response should contain an image");
321
322        assert_eq!(parsed.image, b"final image");
323    }
324
325    #[test]
326    fn response_parsing_rejects_text_only_response() {
327        let response = GenerateContentResponse {
328            candidates: vec![ContentCandidate {
329                content: Some(Content {
330                    role: Some(Role::Model),
331                    parts: vec![Part {
332                        thought: Some(false),
333                        thought_signature: None,
334                        part: PartKind::Text("No image".to_string()),
335                        additional_params: None,
336                    }],
337                }),
338                finish_reason: Some(FinishReason::Stop),
339                safety_ratings: None,
340                citation_metadata: None,
341                token_count: None,
342                avg_logprobs: None,
343                logprobs_result: None,
344                index: None,
345                finish_message: None,
346            }],
347            prompt_feedback: None,
348            usage_metadata: None,
349            model_version: Some(GEMINI_2_5_FLASH_IMAGE.to_string()),
350            response_id: "response-id".to_string(),
351        };
352
353        let err = image_generation::ImageGenerationResponse::<GenerateContentResponse>::try_from(
354            response,
355        )
356        .expect_err("text-only responses should fail");
357
358        assert!(err.to_string().contains("did not include image data"));
359    }
360
361    #[test]
362    fn api_response_parsing_keeps_blocked_prompt_as_success() {
363        let response: ApiResponse<GenerateContentResponse> = serde_json::from_value(json!({
364            "promptFeedback": {
365                "blockReason": "SAFETY"
366            }
367        }))
368        .expect("blocked prompt response should deserialize");
369
370        match response {
371            ApiResponse::Ok(response) => assert!(response.candidates.is_empty()),
372            ApiResponse::Err(err) => panic!("expected success envelope, got error: {err:?}"),
373        }
374    }
375
376    #[tokio::test]
377    async fn image_generation_non_success_preserves_status_and_body() {
378        use crate::client::image_generation::ImageGenerationClient;
379        use crate::image_generation::ImageGenerationModel as _;
380        use crate::test_utils::RecordingHttpClient;
381
382        let body = r#"{"error":{"code":503,"message":"boom","status":"UNAVAILABLE"}}"#;
383        let http_client =
384            RecordingHttpClient::with_error_response(http::StatusCode::SERVICE_UNAVAILABLE, body);
385        let client = Client::builder()
386            .api_key("test-key")
387            .http_client(http_client)
388            .build()
389            .expect("build client");
390        let model = client.image_generation_model(GEMINI_2_5_FLASH_IMAGE);
391
392        let error = model
393            .image_generation(image_generation_request("draw a cat"))
394            .await
395            .expect_err("should fail with non-success status");
396
397        assert!(matches!(error, ImageGenerationError::HttpError(_)));
398        assert_eq!(
399            error.provider_response_status(),
400            Some(http::StatusCode::SERVICE_UNAVAILABLE)
401        );
402        assert_eq!(error.provider_response_body(), Some(body));
403    }
404
405    #[tokio::test]
406    async fn image_generation_2xx_error_envelope_preserves_status_and_body() {
407        use crate::client::image_generation::ImageGenerationClient;
408        use crate::image_generation::ImageGenerationModel as _;
409        use crate::test_utils::RecordingHttpClient;
410
411        // 200 OK carrying Gemini's standard nested error envelope. The error
412        // variant must be tried first because all identifying fields in
413        // `GenerateContentResponse` can be omitted.
414        let body = r#"{"error":{"code":503,"message":"boom","status":"UNAVAILABLE"}}"#;
415        let http_client = RecordingHttpClient::new(body); // 200 OK
416        let client = Client::builder()
417            .api_key("test-key")
418            .http_client(http_client)
419            .build()
420            .expect("build client");
421        let model = client.image_generation_model(GEMINI_2_5_FLASH_IMAGE);
422
423        let error = model
424            .image_generation(image_generation_request("draw a cat"))
425            .await
426            .expect_err("should fail with provider error envelope");
427
428        match &error {
429            ImageGenerationError::ProviderResponse(stored) => {
430                assert_eq!(stored.body, body);
431                assert_eq!(stored.status, Some(http::StatusCode::OK));
432            }
433            other => panic!("expected ProviderResponse, got {other:?}"),
434        }
435    }
436}