use serde::{Deserialize, Serialize};
use url::Url;
use crate::{
errors::OapiError,
images::{Background, ImageResponseFormat, OutputFormat},
rest::post::{Post, PostNoStream},
};
#[derive(Debug, Serialize, Deserialize, Default, Clone)]
pub struct ImageGenerateRequest {
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 moderation: Option<Moderation>,
#[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 style: Option<Style>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
}
#[derive(Debug, Serialize, Deserialize, Clone, Copy)]
#[serde(rename_all = "snake_case")]
pub enum Moderation {
Low,
Auto,
}
#[derive(Debug, Serialize, Deserialize, Clone, Copy)]
#[serde(rename_all = "snake_case")]
pub enum Quality {
Standard,
Hd,
Low,
Medium,
High,
Auto,
}
#[derive(Debug, Serialize, Deserialize, Clone, Copy)]
#[serde(rename_all = "snake_case")]
pub enum Style {
Vivid,
Natural,
}
impl Post for ImageGenerateRequest {
#[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("generations");
Ok(url.to_string())
}
}
impl PostNoStream for ImageGenerateRequest {
type Response = crate::images::ImagesResponse;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dall_e_3_serialization() {
let request = ImageGenerateRequest {
prompt: "A nebula painted in watercolor".to_string(),
model: Some("dall-e-3".to_string()),
n: Some(1),
quality: Some(Quality::Hd),
response_format: Some(ImageResponseFormat::Url),
size: Some("1792x1024".to_string()),
style: Some(Style::Vivid),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(
json.contains(r#""prompt":"A nebula painted in watercolor""#),
"json: {json}"
);
assert!(json.contains(r#""model":"dall-e-3""#), "json: {json}");
assert!(json.contains(r#""quality":"hd""#), "json: {json}");
assert!(json.contains(r#""response_format":"url""#), "json: {json}");
assert!(json.contains(r#""size":"1792x1024""#), "json: {json}");
assert!(json.contains(r#""style":"vivid""#), "json: {json}");
assert!(!json.contains(r#""background""#), "json: {json}");
assert!(!json.contains(r#""moderation""#), "json: {json}");
assert!(!json.contains(r#""user""#), "json: {json}");
}
#[test]
fn gpt_image_serialization() {
let request = ImageGenerateRequest {
prompt: "A painted nebula".to_string(),
model: Some("gpt-image-1".to_string()),
background: Some(Background::Transparent),
moderation: Some(Moderation::Low),
output_compression: Some(80),
output_format: Some(OutputFormat::Webp),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(
json.contains(r#""background":"transparent""#),
"json: {json}"
);
assert!(json.contains(r#""moderation":"low""#), "json: {json}");
assert!(json.contains(r#""output_compression":80"#), "json: {json}");
assert!(json.contains(r#""output_format":"webp""#), "json: {json}");
}
#[test]
fn test_build_url() {
let request = ImageGenerateRequest::default();
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/images/generations");
}
}