gproxy_transform/transform/images/create/
imagen.rs1use crate::protocol::{gemini, openai};
4use crate::transform::context::report_unsupported;
5use crate::transform::{TransformContext, TransformError};
6
7pub mod openai_to_gemini {
8 use super::*;
9
10 pub fn request(
11 input: openai::ImageGenerationRequest,
12 _: &TransformContext,
13 ) -> Result<gemini::ImagenPredictRequest, TransformError> {
14 if input.background.is_some() {
15 report_unsupported("background", "Imagen predict has no background control");
16 }
17 let output_options = input.output_format.as_ref().map(|format| {
18 crate::protocol::wire!(gemini::ImagenOutputOptions {
19 mime_type: Some(format!("image/{}", wire_str(format))),
20 compression_quality: input.output_compression,
21 extra: Default::default(),
22 })
23 });
24 Ok(crate::protocol::wire!(gemini::ImagenPredictRequest {
25 instances: vec![crate::protocol::wire!(gemini::ImagenInstance {
26 prompt: Some(input.prompt),
27 extra: Default::default(),
28 })],
29 parameters: Some(crate::protocol::wire!(gemini::ImagenParameters {
30 sample_count: input.n,
31 aspect_ratio: input.size.as_ref().and_then(size_to_aspect),
32 image_size: None,
33 person_generation: None,
34 negative_prompt: None,
35 seed: None,
36 output_options,
37 extra: Default::default(),
38 })),
39 extra: Default::default(),
40 }))
41 }
42
43 pub fn response(
44 input: openai::ImagesResponse,
45 _: &TransformContext,
46 ) -> Result<gemini::ImagenPredictResponse, TransformError> {
47 let predictions = input
48 .data
49 .into_iter()
50 .flatten()
51 .filter_map(|image| {
52 if image.b64_json.is_none() {
53 report_unsupported(
54 "data[].url",
55 "Imagen predictions carry inline bytes; URL-only images cannot be converted",
56 );
57 return None;
58 }
59 Some(crate::protocol::wire!(gemini::ImagenPrediction {
60 bytes_base64_encoded: image.b64_json,
61 mime_type: input
62 .output_format
63 .as_ref()
64 .map(|format| format!("image/{}", wire_str(format))),
65 rai_filtered_reason: None,
66 extra: Default::default(),
67 }))
68 })
69 .collect();
70 Ok(crate::protocol::wire!(gemini::ImagenPredictResponse {
71 predictions,
72 extra: Default::default(),
73 }))
74 }
75}
76
77pub mod gemini_to_openai {
78 use super::*;
79
80 pub fn request(
81 input: gemini::ImagenPredictRequest,
82 _: &TransformContext,
83 ) -> Result<openai::ImageGenerationRequest, TransformError> {
84 let prompt = input
85 .instances
86 .into_iter()
87 .next()
88 .and_then(|instance| instance.prompt)
89 .ok_or_else(|| TransformError::InvalidInput {
90 reason: "Imagen request has no prompt instance".to_owned(),
91 })?;
92 let parameters = input.parameters.unwrap_or_default();
93 if parameters.negative_prompt.is_some() {
94 report_unsupported(
95 "parameters.negativePrompt",
96 "OpenAI image generation has no negative prompt",
97 );
98 }
99 Ok(crate::protocol::wire!(openai::ImageGenerationRequest {
100 prompt,
101 background: None,
102 model: None,
103 moderation: None,
104 n: parameters.sample_count,
105 output_compression: parameters
106 .output_options
107 .as_ref()
108 .and_then(|options| options.compression_quality),
109 output_format: None,
110 partial_images: None,
111 quality: None,
112 response_format: None,
113 size: aspect_to_size(parameters.aspect_ratio.as_deref()),
114 stream: None,
115 style: None,
116 user: None,
117 extra: Default::default(),
118 }))
119 }
120
121 pub fn response(
122 input: gemini::ImagenPredictResponse,
123 _: &TransformContext,
124 ) -> Result<openai::ImagesResponse, TransformError> {
125 let data = input
126 .predictions
127 .into_iter()
128 .map(|prediction| {
129 crate::protocol::wire!(openai::Image {
130 b64_json: prediction.bytes_base64_encoded,
131 revised_prompt: None,
132 url: None,
133 extra: Default::default(),
134 })
135 })
136 .collect::<Vec<_>>();
137 Ok(crate::protocol::wire!(openai::ImagesResponse {
138 created: 0,
139 background: None,
140 data: Some(data),
141 output_format: None,
142 quality: None,
143 size: None,
144 usage: None,
145 extra: Default::default(),
146 }))
147 }
148}
149
150fn wire_str(value: &impl serde::Serialize) -> String {
151 serde_json::to_value(value)
152 .ok()
153 .and_then(|value| value.as_str().map(str::to_owned))
154 .unwrap_or_default()
155}
156
157fn size_to_aspect(size: &openai::ImageSize) -> Option<String> {
159 let value = wire_str(size);
160 let (width, height) = value.split_once('x')?;
161 let (width, height): (f64, f64) = (width.parse().ok()?, height.parse().ok()?);
162 let ratio = width / height;
163 Some(
164 if ratio > 1.55 {
165 "16:9"
166 } else if ratio > 1.05 {
167 "4:3"
168 } else if ratio > 0.95 {
169 "1:1"
170 } else if ratio > 0.65 {
171 "3:4"
172 } else {
173 "9:16"
174 }
175 .to_owned(),
176 )
177}
178
179fn aspect_to_size(aspect: Option<&str>) -> Option<openai::ImageSize> {
180 let value = match aspect? {
181 "1:1" => "1024x1024",
182 "16:9" | "4:3" => "1536x1024",
183 "9:16" | "3:4" => "1024x1536",
184 _ => return None,
185 };
186 Some(
187 serde_json::from_value(serde_json::json!(value))
188 .unwrap_or(openai::ImageSize::Unknown(value.to_owned())),
189 )
190}
191
192#[cfg(test)]
193mod tests {
194 use serde_json::json;
195
196 use super::*;
197 use crate::protocol::{Operation, OperationKey, Provider};
198
199 #[test]
201 fn maps_native_imagen_request_and_response() {
202 let ctx = TransformContext::new(
203 OperationKey::provider(Operation::CreateImage, Provider::OpenAi),
204 OperationKey::provider(Operation::CreateImage, Provider::Gemini),
205 );
206 let request: openai::ImageGenerationRequest = serde_json::from_value(json!({
207 "prompt": "远山水墨",
208 "model": "imagen-4.0-generate-001",
209 "n": 2,
210 "size": "1024x1536"
211 }))
212 .unwrap();
213 let predict = openai_to_gemini::request(request, &ctx).unwrap();
214 let value = serde_json::to_value(&predict).unwrap();
215 assert_eq!(value["instances"][0]["prompt"], "远山水墨");
216 assert_eq!(value["parameters"]["sampleCount"], 2);
217 assert_eq!(value["parameters"]["aspectRatio"], "3:4");
218
219 let response: gemini::ImagenPredictResponse = serde_json::from_value(json!({
220 "predictions": [{"bytesBase64Encoded": "AAAA", "mimeType": "image/png"}]
221 }))
222 .unwrap();
223 let images = gemini_to_openai::response(response, &ctx).unwrap();
224 assert_eq!(images.data.as_ref().unwrap().len(), 1);
225 assert_eq!(images.data.unwrap()[0].b64_json.as_deref(), Some("AAAA"));
226 }
227}