1use 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
18pub struct OpenAiImageProvider {
20 client: reqwest::Client,
21 base_url: String,
22 api_key: String,
23 store: Arc<dyn ImageSink>,
24}
25
26impl OpenAiImageProvider {
27 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
117fn 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 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 assert!(sink.saved.lock().is_empty());
191 }
192
193 #[test]
194 fn b64_response_is_decoded_and_persisted() {
195 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}