Skip to main content

rig_core/providers/gemini/
image_generation.rs

1//! Image generation through Gemini's `generateContent` endpoint.
2//!
3//! ```no_run
4//! use rig_core::providers::gemini::{Gemini, image_generation::GEMINI_2_5_FLASH_IMAGE};
5//!
6//! # fn main() -> Result<(), Box<dyn std::error::Error>> {
7//! let model = Gemini::from_env()?.image_generation(GEMINI_2_5_FLASH_IMAGE);
8//! # Ok(())
9//! # }
10//! ```
11
12use super::completion::gemini_api_types::{
13    Content, GenerateContentRequest, GenerateContentResponse, GenerationConfig, ImageConfig, Part,
14    PartKind, ResponseModality, Role,
15};
16use crate::completion::Usage;
17use crate::error::EncodeError;
18use crate::error::ProviderError;
19use crate::image_generation;
20use crate::image_generation::{ImageGenerationRequest, NormalizeImageGenerationResponse};
21use crate::operation::ImageGeneration;
22use crate::providers::internal::wire::classify_marker_keyed_frame;
23use crate::wire::Flow;
24use crate::wire::{
25    Body, Decoder, Descriptor, Encoded, Framing, Mode, Out, Wire, WireEvent, WireFrame,
26};
27use base64::Engine;
28use base64::prelude::BASE64_STANDARD;
29use serde_json::Value;
30
31/// `gemini-2.5-flash-image` image generation model, commonly referred to as Nano Banana.
32pub const GEMINI_2_5_FLASH_IMAGE: &str = super::completion::GEMINI_2_5_FLASH_IMAGE;
33
34impl NormalizeImageGenerationResponse for GenerateContentResponse {
35    fn normalize(self) -> Result<image_generation::ImageGenerationResponse, ProviderError> {
36        let image = first_image_bytes(&self)?;
37        let usage = self
38            .usage_metadata
39            .as_ref()
40            .map(Usage::from)
41            .unwrap_or_default();
42
43        Ok(image_generation::ImageGenerationResponse {
44            model: self.model_version,
45            response_id: Some(self.response_id),
46            usage,
47            ..image_generation::ImageGenerationResponse::new(image)
48        })
49    }
50}
51
52fn generate_content_path(model: &str) -> String {
53    format!("/v1beta/models/{model}:generateContent")
54}
55
56fn create_request_body(generation_request: ImageGenerationRequest) -> Result<Value, EncodeError> {
57    let request = GenerateContentRequest {
58        contents: vec![Content {
59            role: Some(Role::User),
60            parts: vec![Part {
61                thought: None,
62                thought_signature: None,
63                part: PartKind::Text(generation_request.prompt),
64                additional_params: None,
65            }],
66        }],
67        tools: None,
68        tool_config: None,
69        generation_config: Some(GenerationConfig {
70            response_modalities: Some(vec![ResponseModality::Image]),
71            image_config: Some(ImageConfig {
72                aspect_ratio: aspect_ratio(generation_request.width, generation_request.height),
73                image_size: None,
74            }),
75            ..Default::default()
76        }),
77        safety_settings: None,
78        system_instruction: None,
79        cached_content: None,
80        additional_params: None,
81    };
82
83    let mut body = serde_json::to_value(request)?;
84
85    if let Some(additional_params) = generation_request.additional_params {
86        merge_json_deep(&mut body, additional_params);
87    }
88
89    Ok(body)
90}
91
92fn merge_json_deep(target: &mut Value, source: Value) {
93    match (target, source) {
94        (Value::Object(target), Value::Object(source)) => {
95            for (key, value) in source {
96                if let Some(existing) = target.get_mut(&key) {
97                    merge_json_deep(existing, value);
98                } else {
99                    target.insert(key, value);
100                }
101            }
102        }
103        (target, source) => *target = source,
104    }
105}
106
107fn aspect_ratio(width: u32, height: u32) -> Option<String> {
108    match (width, height) {
109        (0, _) | (_, 0) => None,
110        (w, h) if w == h => Some("1:1".to_string()),
111        (w, h) if w.saturating_mul(3) == h.saturating_mul(4) => Some("3:4".to_string()),
112        (w, h) if w.saturating_mul(4) == h.saturating_mul(3) => Some("4:3".to_string()),
113        (w, h) if w.saturating_mul(9) == h.saturating_mul(16) => Some("9:16".to_string()),
114        (w, h) if w.saturating_mul(16) == h.saturating_mul(9) => Some("16:9".to_string()),
115        _ => None,
116    }
117}
118
119fn first_image_bytes(response: &GenerateContentResponse) -> Result<Vec<u8>, ProviderError> {
120    for candidate in &response.candidates {
121        let Some(content) = &candidate.content else {
122            continue;
123        };
124
125        for part in &content.parts {
126            if part.thought == Some(true) {
127                continue;
128            }
129
130            if let PartKind::InlineData(inline_data) = &part.part {
131                if !inline_data.mime_type.starts_with("image/") {
132                    continue;
133                }
134
135                return BASE64_STANDARD.decode(&inline_data.data).map_err(|err| {
136                    ProviderError::Response(format!(
137                        "Gemini image data was not valid base64: {err}"
138                    ))
139                });
140            }
141        }
142    }
143
144    Err(ProviderError::Response(
145        "Gemini image generation response did not include image data".into(),
146    ))
147}
148
149/// The image generation wire: `POST /v1beta/models/{model}:generateContent`.
150///
151/// Both [`Mode`]s request image output and decode a whole response document.
152#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
153pub struct Images {
154    /// The provider this wire speaks to.
155    pub provider: super::GeminiConfig,
156    /// Name of the model, for example [`GEMINI_2_5_FLASH_IMAGE`].
157    pub model: String,
158}
159
160impl Images {
161    /// The image generation wire for `model`.
162    pub fn new(provider: super::GeminiConfig, model: impl Into<String>) -> Self {
163        Self {
164            provider,
165            model: model.into(),
166        }
167    }
168}
169
170impl Wire for Images {
171    type Op = ImageGeneration;
172    type Payload = crate::wire::Encoded;
173    type Frame = crate::wire::WireFrame;
174    type Decoder<'id> = ImagesDecoder;
175
176    fn describe(&self) -> Descriptor<'_> {
177        Descriptor::new(super::PROVIDER_NAME).model(self.model.as_str())
178    }
179
180    fn encode(&self, request: ImageGenerationRequest, _mode: Mode) -> Result<Encoded, EncodeError> {
181        let body = serde_json::to_vec(&create_request_body(request)?)?;
182        let request = http::Request::post(format!(
183            "{}{}?key={}",
184            self.provider.base_url,
185            generate_content_path(&self.model),
186            self.provider.api_key.expose()
187        ))
188        .header(http::header::CONTENT_TYPE, "application/json")
189        .body(Body::Bytes(body))?;
190        // Gemini reports no transport request-id header.
191        Ok(Encoded::new(request, Framing::Whole))
192    }
193
194    fn decoder<'id>(&self) -> Self::Decoder<'id> {
195        ImagesDecoder
196    }
197}
198
199/// Decode the first non-thought image in a `generateContent` reply.
200/// Missing image data and invalid base64 produce response errors.
201#[derive(Default)]
202pub struct ImagesDecoder;
203
204impl<'id> Decoder<'id, ImageGeneration> for ImagesDecoder {
205    type Event = GenerateContentResponse;
206
207    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
208        classify_marker_keyed_frame(
209            &frame.as_str(),
210            &["candidates", "promptFeedback", "usageMetadata"],
211        )
212    }
213
214    fn decode(
215        &mut self,
216        event: Self::Event,
217        out: Out<'id, ImageGeneration>,
218    ) -> Result<Flow, ProviderError> {
219        Ok(out.end(event.normalize()?))
220    }
221}
222
223#[cfg(test)]
224mod tests;