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::usage_of;
13use crate::error::{EncodeError, ProviderError};
14use crate::image_generation;
15use crate::image_generation::ImageGenerationRequest;
16use crate::json_utils::Lenient;
17use crate::operation::ImageGeneration;
18use crate::providers::internal::wire::classify_marker_keyed_frame;
19use crate::wire::{
20    Body, Decoder, Descriptor, Encoded, Flow, Framing, Mode, Out, Wire, WireEvent, WireFrame,
21};
22use base64::Engine;
23use base64::prelude::BASE64_STANDARD;
24use serde_json::{Value, json};
25
26/// `gemini-2.5-flash-image` image generation model, commonly referred to as Nano Banana.
27pub const GEMINI_2_5_FLASH_IMAGE: &str = super::completion::GEMINI_2_5_FLASH_IMAGE;
28
29/// The first non-thought image in `reply`, with its usage and identity.
30///
31/// # Errors
32///
33/// When `reply` holds no image data, or data that is not base64.
34pub fn image_of(reply: &Value) -> Result<image_generation::ImageGenerationResponse, ProviderError> {
35    let data = reply
36        .arr("candidates")
37        .iter()
38        .flat_map(|candidate| {
39            candidate
40                .get("content")
41                .map(|content| content.arr("parts"))
42                .unwrap_or_default()
43        })
44        .filter(|part| part.bool("thought") != Some(true))
45        .filter_map(|part| part.get("inlineData"))
46        .find(|blob| {
47            blob.str("mimeType")
48                .is_some_and(|mime| mime.starts_with("image/"))
49        })
50        .and_then(|blob| blob.str("data"))
51        .ok_or_else(|| {
52            ProviderError::Response(
53                "Gemini image generation response did not include image data".into(),
54            )
55        })?;
56    let image = BASE64_STANDARD.decode(data).map_err(|err| {
57        ProviderError::Response(format!("Gemini image data was not valid base64: {err}"))
58    })?;
59    Ok(image_generation::ImageGenerationResponse {
60        model: reply.str("modelVersion").map(str::to_owned),
61        response_id: Some(reply.str("responseId").unwrap_or_default().to_owned()),
62        usage: reply.get("usageMetadata").map(usage_of).unwrap_or_default(),
63        ..image_generation::ImageGenerationResponse::new(image)
64    })
65}
66
67fn create_request_body(generation_request: ImageGenerationRequest) -> Value {
68    let mut image_config = serde_json::Map::new();
69    if let Some(ratio) = aspect_ratio(generation_request.width, generation_request.height) {
70        image_config.insert("aspectRatio".to_owned(), json!(ratio));
71    }
72    let mut body = json!({
73        "contents": [{ "role": "user", "parts": [{ "text": generation_request.prompt }] }],
74        "toolConfig": null,
75        "generationConfig": { "responseModalities": ["IMAGE"], "imageConfig": image_config },
76        "safetySettings": null,
77        "systemInstruction": null,
78    });
79    if let Some(additional_params) = generation_request.additional_params {
80        merge_json_deep(&mut body, additional_params);
81    }
82    body
83}
84
85fn merge_json_deep(target: &mut Value, source: Value) {
86    match (target, source) {
87        (Value::Object(target), Value::Object(source)) => {
88            for (key, value) in source {
89                if let Some(existing) = target.get_mut(&key) {
90                    merge_json_deep(existing, value);
91                } else {
92                    target.insert(key, value);
93                }
94            }
95        }
96        (target, source) => *target = source,
97    }
98}
99
100fn aspect_ratio(width: u32, height: u32) -> Option<String> {
101    match (width, height) {
102        (0, _) | (_, 0) => None,
103        (w, h) if w == h => Some("1:1".to_string()),
104        (w, h) if w.saturating_mul(3) == h.saturating_mul(4) => Some("3:4".to_string()),
105        (w, h) if w.saturating_mul(4) == h.saturating_mul(3) => Some("4:3".to_string()),
106        (w, h) if w.saturating_mul(9) == h.saturating_mul(16) => Some("9:16".to_string()),
107        (w, h) if w.saturating_mul(16) == h.saturating_mul(9) => Some("16:9".to_string()),
108        _ => None,
109    }
110}
111
112/// The image generation wire: `POST /v1beta/models/{model}:generateContent`.
113///
114/// Both [`Mode`]s request image output and decode a whole response document.
115#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
116pub struct Images {
117    /// The provider this wire speaks to.
118    pub provider: super::GeminiConfig,
119    /// Name of the model, for example [`GEMINI_2_5_FLASH_IMAGE`].
120    pub model: String,
121}
122
123impl Images {
124    /// The image generation wire for `model`.
125    pub fn new(provider: super::GeminiConfig, model: impl Into<String>) -> Self {
126        Self {
127            provider,
128            model: model.into(),
129        }
130    }
131}
132
133impl Wire for Images {
134    type Op = ImageGeneration;
135    type Payload = crate::wire::Encoded;
136    type Frame = crate::wire::WireFrame;
137    type Decoder<'id> = ImagesDecoder;
138    type Reassembler = crate::wire::document::Unreassembled;
139
140    fn describe(&self) -> Descriptor<'_> {
141        Descriptor::new(super::PROVIDER_NAME).model(self.model.as_str())
142    }
143
144    fn encode(&self, request: ImageGenerationRequest, _mode: Mode) -> Result<Encoded, EncodeError> {
145        let body = serde_json::to_vec(&create_request_body(request))?;
146        let request = http::Request::post(format!(
147            "{}/v1beta/models/{}:generateContent?key={}",
148            self.provider.base_url,
149            self.model,
150            self.provider.api_key.expose()
151        ))
152        .header(http::header::CONTENT_TYPE, "application/json")
153        .body(Body::Bytes(body))?;
154        // Gemini reports no transport request-id header.
155        Ok(Encoded::new(request, Framing::Whole))
156    }
157
158    fn decoder<'id>(&self) -> Self::Decoder<'id> {
159        ImagesDecoder
160    }
161}
162
163/// Decode the first non-thought image in a `generateContent` reply.
164/// Missing image data and invalid base64 produce response errors.
165#[derive(Default)]
166pub struct ImagesDecoder;
167
168impl<'id> Decoder<'id, ImageGeneration> for ImagesDecoder {
169    type Event = Value;
170
171    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
172        classify_marker_keyed_frame(
173            &frame.as_str(),
174            &["candidates", "promptFeedback", "usageMetadata"],
175        )
176    }
177
178    fn decode(
179        &mut self,
180        event: Self::Event,
181        out: Out<'id, ImageGeneration>,
182    ) -> Result<Flow, ProviderError> {
183        Ok(out.end(image_of(&event)?))
184    }
185}
186
187impl super::GeminiConfig {
188    /// The image generation wire.
189    pub(crate) fn image_generation(&self, model: impl Into<String>) -> Images {
190        Images::new(self.clone(), model)
191    }
192}
193
194#[cfg(test)]
195mod tests;