Skip to main content

codex_api/endpoint/
images.rs

1use crate::auth::SharedAuthProvider;
2use crate::endpoint::session::EndpointSession;
3use crate::error::ApiError;
4use crate::images::ImageEditRequest;
5use crate::images::ImageGenerationRequest;
6use crate::images::ImageResponse;
7use crate::provider::Provider;
8use codex_client::HttpTransport;
9use codex_client::RequestTelemetry;
10use http::HeaderMap;
11use http::Method;
12use serde::Serialize;
13use serde_json::to_value;
14use std::sync::Arc;
15
16pub struct ImagesClient<T: HttpTransport> {
17    session: EndpointSession<T>,
18}
19
20impl<T: HttpTransport> ImagesClient<T> {
21    pub fn new(transport: T, provider: Provider, auth: SharedAuthProvider) -> Self {
22        Self {
23            session: EndpointSession::new(transport, provider, auth),
24        }
25    }
26
27    pub fn with_telemetry(self, request: Option<Arc<dyn RequestTelemetry>>) -> Self {
28        Self {
29            session: self.session.with_request_telemetry(request),
30        }
31    }
32
33    pub async fn generate(
34        &self,
35        request: &ImageGenerationRequest,
36        extra_headers: HeaderMap,
37    ) -> Result<ImageResponse, ApiError> {
38        self.post_image_request(
39            "images/generations",
40            request,
41            extra_headers,
42            "image generation",
43        )
44        .await
45    }
46
47    pub async fn edit(
48        &self,
49        request: &ImageEditRequest,
50        extra_headers: HeaderMap,
51    ) -> Result<ImageResponse, ApiError> {
52        self.post_image_request("images/edits", request, extra_headers, "image edit")
53            .await
54    }
55
56    async fn post_image_request<R: Serialize>(
57        &self,
58        path: &str,
59        request: &R,
60        extra_headers: HeaderMap,
61        operation: &str,
62    ) -> Result<ImageResponse, ApiError> {
63        let body = to_value(request)
64            .map_err(|e| ApiError::Stream(format!("failed to encode {operation} request: {e}")))?;
65        let resp = self
66            .session
67            .execute(Method::POST, path, extra_headers, Some(body))
68            .await?;
69        serde_json::from_slice(&resp.body)
70            .map_err(|e| ApiError::Stream(format!("failed to decode {operation} response: {e}")))
71    }
72}
73
74#[cfg(test)]
75mod tests {
76    use super::*;
77    use crate::auth::AuthProvider;
78    use crate::images::ImageBackground;
79    use crate::images::ImageData;
80    use crate::images::ImageQuality;
81    use crate::images::ImageUrl;
82    use crate::provider::RetryConfig;
83    use codex_client::Request;
84    use codex_client::RequestBody;
85    use codex_client::Response;
86    use codex_client::StreamResponse;
87    use codex_client::TransportError;
88    use http::StatusCode;
89    use pretty_assertions::assert_eq;
90    use serde_json::json;
91    use std::sync::Mutex;
92    use std::time::Duration;
93
94    #[derive(Clone, Default)]
95    struct DummyAuth;
96
97    impl AuthProvider for DummyAuth {
98        fn add_auth_headers(&self, _headers: &mut HeaderMap) {}
99    }
100
101    #[derive(Clone)]
102    struct CapturingTransport {
103        last_request: Arc<Mutex<Option<Request>>>,
104        response_body: Arc<Vec<u8>>,
105    }
106
107    impl CapturingTransport {
108        fn new(response_body: Vec<u8>) -> Self {
109            Self {
110                last_request: Arc::new(Mutex::new(None)),
111                response_body: Arc::new(response_body),
112            }
113        }
114    }
115
116    impl HttpTransport for CapturingTransport {
117        async fn execute(&self, req: Request) -> Result<Response, TransportError> {
118            *self.last_request.lock().expect("lock request store") = Some(req);
119            Ok(Response {
120                status: StatusCode::OK,
121                headers: HeaderMap::new(),
122                body: self.response_body.as_ref().clone().into(),
123            })
124        }
125
126        async fn stream(&self, _req: Request) -> Result<StreamResponse, TransportError> {
127            Err(TransportError::Build("stream should not run".to_string()))
128        }
129    }
130
131    fn provider() -> Provider {
132        Provider {
133            name: "test".to_string(),
134            base_url: "https://example.com/api/codex".to_string(),
135            query_params: None,
136            headers: HeaderMap::new(),
137            retry: RetryConfig {
138                max_attempts: 1,
139                base_delay: Duration::from_millis(1),
140                retry_429: false,
141                retry_5xx: true,
142                retry_transport: true,
143            },
144            stream_idle_timeout: Duration::from_secs(1),
145        }
146    }
147
148    fn response_body() -> Vec<u8> {
149        serde_json::to_vec(&json!({
150            "created": 1778832973u64,
151            "background": "opaque",
152            "data": [{"b64_json": "REDACT"}],
153            "output_format": "png",
154            "quality": "medium",
155            "size": "1024x1536",
156            "usage": {
157                "input_tokens": 1474,
158                "input_tokens_details": {
159                    "image_tokens": 1457,
160                    "text_tokens": 17,
161                },
162                "output_tokens": 1372,
163                "output_tokens_details": {
164                    "image_tokens": 1372,
165                    "text_tokens": 0,
166                },
167                "total_tokens": 2846,
168            }
169        }))
170        .expect("serialize response")
171    }
172
173    fn expected_response() -> ImageResponse {
174        ImageResponse {
175            created: 1778832973,
176            background: Some(ImageBackground::Opaque),
177            data: vec![ImageData {
178                b64_json: "REDACT".to_string(),
179            }],
180            quality: Some(ImageQuality::Medium),
181            size: Some("1024x1536".to_string()),
182        }
183    }
184
185    fn captured_request(transport: &CapturingTransport) -> Request {
186        transport
187            .last_request
188            .lock()
189            .expect("lock request store")
190            .clone()
191            .expect("request should be captured")
192    }
193
194    #[tokio::test]
195    async fn generate_posts_typed_request_and_parses_image_response() {
196        let transport = CapturingTransport::new(response_body());
197        let client = ImagesClient::new(transport.clone(), provider(), Arc::new(DummyAuth));
198
199        let response = client
200            .generate(
201                &ImageGenerationRequest {
202                    prompt: "a red fox in a field".to_string(),
203                    background: Some(ImageBackground::Opaque),
204                    model: "gpt-image-1.5".to_string(),
205                    n: None,
206                    quality: Some(ImageQuality::Medium),
207                    size: Some("1024x1536".to_string()),
208                },
209                HeaderMap::new(),
210            )
211            .await
212            .expect("image generation request should succeed");
213
214        assert_eq!(response, expected_response());
215
216        let request = captured_request(&transport);
217        assert_eq!(
218            request.url,
219            "https://example.com/api/codex/images/generations"
220        );
221        assert_eq!(
222            request.body.as_ref().and_then(RequestBody::json),
223            Some(&json!({
224                "prompt": "a red fox in a field",
225                "background": "opaque",
226                "model": "gpt-image-1.5",
227                "quality": "medium",
228                "size": "1024x1536",
229            }))
230        );
231    }
232
233    #[tokio::test]
234    async fn edit_posts_typed_request_and_parses_image_response() {
235        let transport = CapturingTransport::new(response_body());
236        let client = ImagesClient::new(transport.clone(), provider(), Arc::new(DummyAuth));
237
238        let response = client
239            .edit(
240                &ImageEditRequest {
241                    images: vec![ImageUrl {
242                        image_url: "data:image/png;base64,Zm9v".to_string(),
243                    }],
244                    prompt: "add a red hat".to_string(),
245                    background: None,
246                    model: "gpt-image-1.5".to_string(),
247                    n: None,
248                    quality: None,
249                    size: None,
250                },
251                HeaderMap::new(),
252            )
253            .await
254            .expect("image edit request should succeed");
255
256        assert_eq!(response, expected_response());
257
258        let request = captured_request(&transport);
259        assert_eq!(request.url, "https://example.com/api/codex/images/edits");
260        assert_eq!(
261            request.body.as_ref().and_then(RequestBody::json),
262            Some(&json!({
263                "images": [{"image_url": "data:image/png;base64,Zm9v"}],
264                "prompt": "add a red hat",
265                "model": "gpt-image-1.5",
266            }))
267        );
268    }
269
270    #[tokio::test]
271    async fn image_response_requires_image_data() {
272        let transport = CapturingTransport::new(
273            serde_json::to_vec(&json!({"created": 1778832973u64})).expect("serialize response"),
274        );
275        let client = ImagesClient::new(transport, provider(), Arc::new(DummyAuth));
276
277        let error = client
278            .generate(
279                &ImageGenerationRequest {
280                    prompt: "a red fox in a field".to_string(),
281                    background: None,
282                    model: "gpt-image-1.5".to_string(),
283                    n: None,
284                    quality: None,
285                    size: None,
286                },
287                HeaderMap::new(),
288            )
289            .await
290            .expect_err("image response without data should fail");
291
292        let ApiError::Stream(message) = error else {
293            panic!("expected image response decode error");
294        };
295        assert!(
296            message.starts_with("failed to decode image generation response: missing field `data`"),
297            "{message}"
298        );
299    }
300}