use std::path::PathBuf;
use serde::Serialize;
use url::Url;
use crate::{
errors::OapiError,
images::{Background, ImageResponseFormat, OutputFormat, enum_to_literal},
rest::post::{Post, PostNoStream},
};
#[derive(Debug, Serialize, Default, Clone)]
pub struct ImageEditRequest {
#[serde(skip_serializing)]
pub image: Vec<PathBuf>,
pub prompt: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub background: Option<Background>,
#[serde(skip_serializing_if = "Option::is_none")]
pub input_fidelity: Option<InputFidelity>,
#[serde(skip_serializing)]
pub mask: Option<PathBuf>,
#[serde(skip_serializing_if = "Option::is_none")]
pub n: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output_compression: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output_format: Option<OutputFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub quality: Option<Quality>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<ImageResponseFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub size: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "snake_case")]
pub enum InputFidelity {
High,
Low,
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "snake_case")]
pub enum Quality {
Standard,
Low,
Medium,
High,
Auto,
}
impl Post for ImageEditRequest {
#[inline]
fn is_streaming(&self) -> bool {
false
}
fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
url.path_segments_mut()
.map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
.push("images")
.push("edits");
Ok(url.to_string())
}
}
impl PostNoStream for ImageEditRequest {
type Response = crate::images::ImagesResponse;
async fn get_response_string(
&self,
client: &reqwest::Client,
url: &str,
key: &str,
) -> Result<String, OapiError> {
if self.image.is_empty() {
return Err(OapiError::ResponseError(
"At least one image is required".to_string(),
));
}
let mut form = reqwest::multipart::Form::new();
let image_part_name = if self.image.len() == 1 {
"image"
} else {
"image[]"
};
for path in &self.image {
let content = tokio::fs::read(path).await?;
let file_name = file_name_of(path)?;
let part = reqwest::multipart::Part::bytes(content).file_name(file_name);
form = form.part(image_part_name, part);
}
if let Some(mask) = &self.mask {
let content = tokio::fs::read(mask).await?;
let file_name = file_name_of(mask)?;
let part = reqwest::multipart::Part::bytes(content).file_name(file_name);
form = form.part("mask", part);
}
form = form.text("prompt", self.prompt.clone());
if let Some(model) = &self.model {
form = form.text("model", model.clone());
}
if let Some(background) = self.background {
form = form.text("background", enum_to_literal(&background)?);
}
if let Some(input_fidelity) = self.input_fidelity {
form = form.text("input_fidelity", enum_to_literal(&input_fidelity)?);
}
if let Some(n) = self.n {
form = form.text("n", n.to_string());
}
if let Some(output_compression) = self.output_compression {
form = form.text("output_compression", output_compression.to_string());
}
if let Some(output_format) = self.output_format {
form = form.text("output_format", enum_to_literal(&output_format)?);
}
if let Some(quality) = self.quality {
form = form.text("quality", enum_to_literal(&quality)?);
}
if let Some(response_format) = self.response_format {
form = form.text("response_format", enum_to_literal(&response_format)?);
}
if let Some(size) = &self.size {
form = form.text("size", size.clone());
}
if let Some(user) = &self.user {
form = form.text("user", user.clone());
}
let response = client
.post(url)
.header("Accept", "application/json")
.bearer_auth(key)
.multipart(form)
.send()
.await?;
crate::rest::response_text_checked(response).await
}
}
fn file_name_of(path: &std::path::Path) -> Result<String, OapiError> {
path.file_name()
.and_then(|name| name.to_str())
.map(|name| name.to_string())
.ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_build_url() {
let request = ImageEditRequest::default();
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/images/edits");
}
#[test]
fn enum_literals() {
assert_eq!(enum_to_literal(&InputFidelity::High).unwrap(), "high");
assert_eq!(enum_to_literal(&Quality::Standard).unwrap(), "standard");
assert_eq!(
enum_to_literal(&ImageResponseFormat::B64Json).unwrap(),
"b64_json"
);
}
}