rig_core/providers/gemini/
image_generation.rs1use 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
31pub 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#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
153pub struct Images {
154 pub provider: super::GeminiConfig,
156 pub model: String,
158}
159
160impl Images {
161 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 Ok(Encoded::new(request, Framing::Whole))
192 }
193
194 fn decoder<'id>(&self) -> Self::Decoder<'id> {
195 ImagesDecoder
196 }
197}
198
199#[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;