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