Skip to main content

gproxy_transform/transform/images/create/
imagen.rs

1//! OpenAI create-image <-> Imagen native `:predict` conversions.
2
3use 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
157/// `WxH` -> Imagen aspect ratio (exact pixel dims are provider-fixed).
158fn 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    /// 尺寸→宽高比与 predictions→data 是仅有的非平凡映射。
200    #[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}