Skip to main content

oxios_kernel/image_gen/
openai.rs

1//! OpenAI `/v1/images/generations` provider.
2//!
3//! Also works with OpenAI-compatible endpoints (local Stable Diffusion WebUI,
4//! Azure OpenAI, OpenRouter images, …) given the right `base_url`. Response
5//! normalization (§4.3 of the port design) handles both `url` and `b64_json`
6//! return shapes, since which one a model returns is the user's choice.
7
8use std::sync::Arc;
9
10use async_trait::async_trait;
11use base64::Engine;
12use serde::Deserialize;
13
14use crate::image_gen::{
15    GeneratedImage, ImageGenError, ImageGenProvider, ImageGenRequest, ImageGenResult, ImageSink,
16};
17
18/// OpenAI image-generation client.
19pub struct OpenAiImageProvider {
20    client: reqwest::Client,
21    base_url: String,
22    api_key: String,
23    store: Arc<dyn ImageSink>,
24}
25
26impl OpenAiImageProvider {
27    /// Construct from already-resolved values.
28    ///
29    /// The caller (tool/config layer) resolves the credential and image dir;
30    /// this keeps the module fully unit-testable with a fake [`ImageSink`].
31    pub fn new(
32        base_url: String,
33        api_key: String,
34        store: Arc<dyn ImageSink>,
35    ) -> Result<Self, ImageGenError> {
36        let client = reqwest::Client::builder()
37            .timeout(std::time::Duration::from_secs(120))
38            .connect_timeout(std::time::Duration::from_secs(10))
39            .build()
40            .map_err(|e| ImageGenError::Transport(e.to_string()))?;
41        Ok(Self {
42            client,
43            base_url,
44            api_key,
45            store,
46        })
47    }
48}
49
50#[derive(Deserialize)]
51struct OpenAiResponse {
52    data: Vec<OpenAiImage>,
53}
54
55#[derive(Deserialize)]
56struct OpenAiImage {
57    url: Option<String>,
58    b64_json: Option<String>,
59    revised_prompt: Option<String>,
60}
61
62#[async_trait]
63impl ImageGenProvider for OpenAiImageProvider {
64    fn name(&self) -> &'static str {
65        "openai"
66    }
67
68    async fn generate(&self, req: &ImageGenRequest) -> Result<ImageGenResult, ImageGenError> {
69        let model = req.model.as_deref().ok_or(ImageGenError::MissingModel)?;
70        let body = serde_json::json!({
71            "model": model,
72            "prompt": req.prompt,
73            "n": req.n.clamp(1, 8),
74            "size": req.size.map(|s| s.as_str()).unwrap_or("1024x1024"),
75        });
76        let url = format!("{}/images/generations", self.base_url.trim_end_matches('/'));
77
78        let resp = self
79            .client
80            .post(&url)
81            .bearer_auth(&self.api_key)
82            .json(&body)
83            .send()
84            .await
85            .map_err(|e| ImageGenError::Transport(e.to_string()))?;
86
87        let status = resp.status();
88        if !status.is_success() {
89            let body = resp.text().await.unwrap_or_default();
90            return Err(ImageGenError::Http {
91                status: status.as_u16(),
92                body,
93            });
94        }
95
96        let parsed: OpenAiResponse = resp
97            .json()
98            .await
99            .map_err(|e| ImageGenError::BadResponse(format!("invalid JSON: {e}")))?;
100
101        let revised_prompt = parsed
102            .data
103            .iter()
104            .rev()
105            .find_map(|d| d.revised_prompt.clone());
106        let images = normalize_images(&parsed.data, self.store.as_ref())?;
107
108        Ok(ImageGenResult {
109            images,
110            provider: "openai".into(),
111            model: model.into(),
112            revised_prompt,
113        })
114    }
115}
116
117/// Map provider image entries to fetchable URLs. Pure over the [`ImageSink`].
118///
119/// - `url` present  → used directly (e.g. `dall-e-3`).
120/// - only `b64_json` → decoded, persisted via the sink, served URL returned
121///   (e.g. `gpt-image-1`, which always returns base64).
122fn normalize_images(
123    data: &[OpenAiImage],
124    store: &dyn ImageSink,
125) -> Result<Vec<GeneratedImage>, ImageGenError> {
126    data.iter()
127        .map(|img| match (&img.url, &img.b64_json) {
128            (Some(url), _) => Ok(GeneratedImage {
129                url: url.clone(),
130                width: None,
131                height: None,
132            }),
133            (None, Some(b64)) => {
134                let bytes = base64::engine::general_purpose::STANDARD
135                    .decode(b64)
136                    .map_err(|e| ImageGenError::Base64(e.to_string()))?;
137                let url = store.save(bytes, "png")?;
138                Ok(GeneratedImage {
139                    url,
140                    width: None,
141                    height: None,
142                })
143            }
144            (None, None) => Err(ImageGenError::BadResponse(
145                "image entry has neither url nor b64_json".into(),
146            )),
147        })
148        .collect()
149}
150
151#[cfg(test)]
152mod tests {
153    use parking_lot::Mutex;
154
155    use super::*;
156    use crate::image_gen::ImageSink;
157
158    /// Records saves instead of hitting disk — lets us test `normalize_images`
159    /// without a real provider/HTTP.
160    struct FakeSink {
161        saved: Mutex<Vec<(Vec<u8>, String)>>,
162    }
163
164    impl ImageSink for FakeSink {
165        fn save(&self, bytes: Vec<u8>, ext: &str) -> Result<String, ImageGenError> {
166            let n = self.saved.lock().len();
167            let url = format!("/api/images/fake-{n}.{ext}");
168            self.saved.lock().push((bytes, ext.into()));
169            Ok(url)
170        }
171    }
172
173    fn img(url: Option<&str>, b64: Option<&str>) -> OpenAiImage {
174        OpenAiImage {
175            url: url.map(Into::into),
176            b64_json: b64.map(Into::into),
177            revised_prompt: None,
178        }
179    }
180
181    #[test]
182    fn url_response_passes_through() {
183        let sink = FakeSink {
184            saved: Mutex::new(vec![]),
185        };
186        let out = normalize_images(&[img(Some("https://cdn/x.png"), None)], &sink).unwrap();
187        assert_eq!(out.len(), 1);
188        assert_eq!(out[0].url, "https://cdn/x.png");
189        // Nothing persisted — url used directly.
190        assert!(sink.saved.lock().is_empty());
191    }
192
193    #[test]
194    fn b64_response_is_decoded_and_persisted() {
195        // "hello" → base64
196        let b64 = base64::engine::general_purpose::STANDARD.encode(b"hello");
197        let sink = FakeSink {
198            saved: Mutex::new(vec![]),
199        };
200        let out = normalize_images(&[img(None, Some(&b64))], &sink).unwrap();
201        assert_eq!(out.len(), 1);
202        assert!(out[0].url.starts_with("/api/images/fake-0.png"));
203        let (bytes, ext) = &sink.saved.lock()[0];
204        assert_eq!(bytes, b"hello");
205        assert_eq!(ext, "png");
206    }
207
208    #[test]
209    fn mixed_batch_normalizes_each_entry() {
210        let b64 = base64::engine::general_purpose::STANDARD.encode(b"img");
211        let sink = FakeSink {
212            saved: Mutex::new(vec![]),
213        };
214        let out = normalize_images(
215            &[img(Some("https://cdn/a.png"), None), img(None, Some(&b64))],
216            &sink,
217        )
218        .unwrap();
219        assert_eq!(out.len(), 2);
220        assert_eq!(out[0].url, "https://cdn/a.png");
221        assert!(out[1].url.starts_with("/api/images/"));
222        assert_eq!(sink.saved.lock().len(), 1);
223    }
224
225    #[test]
226    fn missing_url_and_b64_is_an_error() {
227        let sink = FakeSink {
228            saved: Mutex::new(vec![]),
229        };
230        let err = normalize_images(&[img(None, None)], &sink).unwrap_err();
231        assert!(matches!(err, ImageGenError::BadResponse(_)));
232    }
233}