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::post::{Post, PostNoStream},
21};
22
23#[derive(Debug, Serialize, Default, Clone)]
25pub struct ImageEditRequest {
26 #[serde(skip_serializing)]
36 pub image: Vec<PathBuf>,
37 pub prompt: String,
42 #[serde(skip_serializing_if = "Option::is_none")]
47 pub model: Option<String>,
48 #[serde(skip_serializing_if = "Option::is_none")]
52 pub background: Option<Background>,
53 #[serde(skip_serializing_if = "Option::is_none")]
59 pub input_fidelity: Option<InputFidelity>,
60 #[serde(skip_serializing)]
66 pub mask: Option<PathBuf>,
67 #[serde(skip_serializing_if = "Option::is_none")]
69 pub n: Option<u32>,
70 #[serde(skip_serializing_if = "Option::is_none")]
75 pub output_compression: Option<u32>,
76 #[serde(skip_serializing_if = "Option::is_none")]
81 pub output_format: Option<OutputFormat>,
82 #[serde(skip_serializing_if = "Option::is_none")]
85 pub quality: Option<Quality>,
86 #[serde(skip_serializing_if = "Option::is_none")]
91 pub response_format: Option<ImageResponseFormat>,
92 #[serde(skip_serializing_if = "Option::is_none")]
95 pub size: Option<String>,
96 #[serde(skip_serializing_if = "Option::is_none")]
99 pub user: Option<String>,
100}
101
102#[derive(Debug, Serialize, Clone, Copy)]
105#[serde(rename_all = "snake_case")]
106pub enum InputFidelity {
107 High,
108 Low,
109}
110
111#[derive(Debug, Serialize, Clone, Copy)]
113#[serde(rename_all = "snake_case")]
114pub enum Quality {
115 Standard,
116 Low,
117 Medium,
118 High,
119 Auto,
120}
121
122impl Post for ImageEditRequest {
123 #[inline]
124 fn is_streaming(&self) -> bool {
125 false
126 }
127
128 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
132 let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
133 url.path_segments_mut()
134 .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
135 .push("images")
136 .push("edits");
137
138 Ok(url.to_string())
139 }
140}
141
142impl PostNoStream for ImageEditRequest {
143 type Response = crate::images::ImagesResponse;
144
145 async fn get_response_string(
148 &self,
149 client: &reqwest::Client,
150 url: &str,
151 key: &str,
152 ) -> Result<String, OapiError> {
153 if self.image.is_empty() {
154 return Err(OapiError::ResponseError(
155 "At least one image is required".to_string(),
156 ));
157 }
158
159 let mut form = reqwest::multipart::Form::new();
160
161 let image_part_name = if self.image.len() == 1 {
164 "image"
165 } else {
166 "image[]"
167 };
168 for path in &self.image {
169 let content = tokio::fs::read(path).await?;
170 let file_name = file_name_of(path)?;
171 let part = reqwest::multipart::Part::bytes(content).file_name(file_name);
172 form = form.part(image_part_name, part);
173 }
174
175 if let Some(mask) = &self.mask {
176 let content = tokio::fs::read(mask).await?;
177 let file_name = file_name_of(mask)?;
178 let part = reqwest::multipart::Part::bytes(content).file_name(file_name);
179 form = form.part("mask", part);
180 }
181
182 form = form.text("prompt", self.prompt.clone());
183
184 if let Some(model) = &self.model {
185 form = form.text("model", model.clone());
186 }
187 if let Some(background) = self.background {
188 form = form.text("background", enum_to_literal(&background)?);
189 }
190 if let Some(input_fidelity) = self.input_fidelity {
191 form = form.text("input_fidelity", enum_to_literal(&input_fidelity)?);
192 }
193 if let Some(n) = self.n {
194 form = form.text("n", n.to_string());
195 }
196 if let Some(output_compression) = self.output_compression {
197 form = form.text("output_compression", output_compression.to_string());
198 }
199 if let Some(output_format) = self.output_format {
200 form = form.text("output_format", enum_to_literal(&output_format)?);
201 }
202 if let Some(quality) = self.quality {
203 form = form.text("quality", enum_to_literal(&quality)?);
204 }
205 if let Some(response_format) = self.response_format {
206 form = form.text("response_format", enum_to_literal(&response_format)?);
207 }
208 if let Some(size) = &self.size {
209 form = form.text("size", size.clone());
210 }
211 if let Some(user) = &self.user {
212 form = form.text("user", user.clone());
213 }
214
215 let response = client
216 .post(url)
217 .header("Accept", "application/json")
218 .bearer_auth(key)
219 .multipart(form)
220 .send()
221 .await?;
222
223 crate::rest::response_text_checked(response).await
224 }
225}
226
227fn file_name_of(path: &std::path::Path) -> Result<String, OapiError> {
229 path.file_name()
230 .and_then(|name| name.to_str())
231 .map(|name| name.to_string())
232 .ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))
233}
234
235#[cfg(test)]
236mod tests {
237 use super::*;
238
239 #[test]
240 fn test_build_url() {
241 let request = ImageEditRequest::default();
242 let url = request.build_url("https://api.openai.com/v1/").unwrap();
243 assert_eq!(url, "https://api.openai.com/v1/images/edits");
244 }
245
246 #[test]
248 fn enum_literals() {
249 assert_eq!(enum_to_literal(&InputFidelity::High).unwrap(), "high");
250 assert_eq!(enum_to_literal(&Quality::Standard).unwrap(), "standard");
251 assert_eq!(
252 enum_to_literal(&ImageResponseFormat::B64Json).unwrap(),
253 "b64_json"
254 );
255 }
256}