openai_interface/images/
edit.rs1use std::path::PathBuf;
13
14use serde::Serialize;
15use url::Url;
16
17use crate::{
18 errors::OapiError,
19 images::{Background, ImageResponseFormat, OutputFormat, enum_to_literal},
20 rest::RequestOptions,
21 rest::post::{Post, PostNoStream},
22};
23
24#[derive(Debug, Serialize, Default, Clone)]
26pub struct ImageEditRequest {
27 #[serde(skip_serializing)]
37 pub image: Vec<PathBuf>,
38 pub prompt: String,
43 #[serde(skip_serializing_if = "Option::is_none")]
48 pub model: Option<String>,
49 #[serde(skip_serializing_if = "Option::is_none")]
53 pub background: Option<Background>,
54 #[serde(skip_serializing_if = "Option::is_none")]
60 pub input_fidelity: Option<InputFidelity>,
61 #[serde(skip_serializing)]
67 pub mask: Option<PathBuf>,
68 #[serde(skip_serializing_if = "Option::is_none")]
70 pub n: Option<u32>,
71 #[serde(skip_serializing_if = "Option::is_none")]
76 pub output_compression: Option<u32>,
77 #[serde(skip_serializing_if = "Option::is_none")]
82 pub output_format: Option<OutputFormat>,
83 #[serde(skip_serializing_if = "Option::is_none")]
86 pub quality: Option<Quality>,
87 #[serde(skip_serializing_if = "Option::is_none")]
92 pub response_format: Option<ImageResponseFormat>,
93 #[serde(skip_serializing_if = "Option::is_none")]
96 pub size: Option<String>,
97 #[serde(skip_serializing_if = "Option::is_none")]
100 pub user: Option<String>,
101 pub extra_body_map: Option<serde_json::Map<String, serde_json::Value>>,
104}
105
106#[derive(Debug, Serialize, Clone, Copy)]
109#[serde(rename_all = "snake_case")]
110pub enum InputFidelity {
111 High,
112 Low,
113}
114
115#[derive(Debug, Serialize, Clone, Copy)]
117#[serde(rename_all = "snake_case")]
118pub enum Quality {
119 Standard,
120 Low,
121 Medium,
122 High,
123 Auto,
124}
125
126impl Post for ImageEditRequest {
127 #[inline]
128 fn is_streaming(&self) -> bool {
129 false
130 }
131
132 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
136 let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
137 url.path_segments_mut()
138 .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
139 .push("images")
140 .push("edits");
141
142 Ok(url.to_string())
143 }
144}
145
146impl PostNoStream for ImageEditRequest {
147 type Response = crate::images::ImagesResponse;
148
149 async fn get_response_string(
152 &self,
153 client: &reqwest::Client,
154 base_url: &str,
155 options: &RequestOptions,
156 ) -> Result<String, OapiError> {
157 if self.image.is_empty() {
158 return Err(OapiError::ResponseError(
159 "At least one image is required".to_string(),
160 ));
161 }
162
163 let mut form = reqwest::multipart::Form::new();
164
165 let image_part_name = if self.image.len() == 1 {
168 "image"
169 } else {
170 "image[]"
171 };
172 for path in &self.image {
173 let content = tokio::fs::read(path).await?;
174 let file_name = file_name_of(path)?;
175 let part = reqwest::multipart::Part::bytes(content).file_name(file_name);
176 form = form.part(image_part_name, part);
177 }
178
179 if let Some(mask) = &self.mask {
180 let content = tokio::fs::read(mask).await?;
181 let file_name = file_name_of(mask)?;
182 let part = reqwest::multipart::Part::bytes(content).file_name(file_name);
183 form = form.part("mask", part);
184 }
185
186 form = form.text("prompt", self.prompt.clone());
187
188 if let Some(model) = &self.model {
189 form = form.text("model", model.clone());
190 }
191 if let Some(background) = self.background {
192 form = form.text("background", enum_to_literal(&background)?);
193 }
194 if let Some(input_fidelity) = self.input_fidelity {
195 form = form.text("input_fidelity", enum_to_literal(&input_fidelity)?);
196 }
197 if let Some(n) = self.n {
198 form = form.text("n", n.to_string());
199 }
200 if let Some(output_compression) = self.output_compression {
201 form = form.text("output_compression", output_compression.to_string());
202 }
203 if let Some(output_format) = self.output_format {
204 form = form.text("output_format", enum_to_literal(&output_format)?);
205 }
206 if let Some(quality) = self.quality {
207 form = form.text("quality", enum_to_literal(&quality)?);
208 }
209 if let Some(response_format) = self.response_format {
210 form = form.text("response_format", enum_to_literal(&response_format)?);
211 }
212 if let Some(size) = &self.size {
213 form = form.text("size", size.clone());
214 }
215 if let Some(user) = &self.user {
216 form = form.text("user", user.clone());
217 }
218
219 form = crate::rest::post::append_extra_body_map(form, &self.extra_body_map);
220
221 let url = self.build_url(base_url)?;
222 crate::rest::post::post_multipart_json(client, url, form, options).await
223 }
224}
225
226fn file_name_of(path: &std::path::Path) -> Result<String, OapiError> {
228 path.file_name()
229 .and_then(|name| name.to_str())
230 .map(|name| name.to_string())
231 .ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))
232}
233
234#[cfg(test)]
235mod tests {
236 use super::*;
237
238 #[test]
239 fn test_build_url() {
240 let request = ImageEditRequest::default();
241 let url = request.build_url("https://api.openai.com/v1/").unwrap();
242 assert_eq!(url, "https://api.openai.com/v1/images/edits");
243 }
244
245 #[test]
247 fn enum_literals() {
248 assert_eq!(enum_to_literal(&InputFidelity::High).unwrap(), "high");
249 assert_eq!(enum_to_literal(&Quality::Standard).unwrap(), "standard");
250 assert_eq!(
251 enum_to_literal(&ImageResponseFormat::B64Json).unwrap(),
252 "b64_json"
253 );
254 }
255}