use std::path::PathBuf;
use serde::Serialize;
use url::Url;
use crate::{
errors::OapiError,
images::{ImageResponseFormat, enum_to_literal},
rest::post::{Post, PostNoStream},
};
#[derive(Debug, Serialize, Default, Clone)]
pub struct ImageVariationRequest {
#[serde(skip_serializing)]
pub image: PathBuf,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub n: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<ImageResponseFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub size: Option<VariationSize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
}
#[derive(Debug, Serialize, Clone, Copy)]
pub enum VariationSize {
#[serde(rename = "256x256")]
S256x256,
#[serde(rename = "512x512")]
S512x512,
#[serde(rename = "1024x1024")]
S1024x1024,
}
impl Post for ImageVariationRequest {
#[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("variations");
Ok(url.to_string())
}
}
impl PostNoStream for ImageVariationRequest {
type Response = crate::images::ImagesResponse;
async fn get_response_string(
&self,
client: &reqwest::Client,
url: &str,
key: &str,
) -> Result<String, OapiError> {
let content = tokio::fs::read(&self.image).await?;
let file_name = self
.image
.file_name()
.and_then(|name| name.to_str())
.ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))?
.to_string();
let image_part = reqwest::multipart::Part::bytes(content).file_name(file_name);
let mut form = reqwest::multipart::Form::new().part("image", image_part);
if let Some(model) = &self.model {
form = form.text("model", model.clone());
}
if let Some(n) = self.n {
form = form.text("n", n.to_string());
}
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", enum_to_literal(&size)?);
}
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
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_build_url() {
let request = ImageVariationRequest::default();
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/images/variations");
}
#[test]
fn size_literals() {
assert_eq!(
enum_to_literal(&VariationSize::S256x256).unwrap(),
"256x256"
);
assert_eq!(
enum_to_literal(&VariationSize::S1024x1024).unwrap(),
"1024x1024"
);
}
}