rig_core/providers/gemini/
image_generation.rs1use 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
26pub const GEMINI_2_5_FLASH_IMAGE: &str = super::completion::GEMINI_2_5_FLASH_IMAGE;
28
29pub 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#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
116pub struct Images {
117 pub provider: super::GeminiConfig,
119 pub model: String,
121}
122
123impl Images {
124 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 Ok(Encoded::new(request, Framing::Whole))
156 }
157
158 fn decoder<'id>(&self) -> Self::Decoder<'id> {
159 ImagesDecoder
160 }
161}
162
163#[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 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;