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}