Skip to main content

openai_interface/images/
variation.rs

1//! Creates a variation of a given image.
2//!
3//! Endpoint: `POST /images/variations` (multipart/form-data request / JSON
4//! response). This endpoint only supports `dall-e-2`.
5//!
6//! > ![warn] This module is untested!
7//! > No OpenAI-compatible provider accessible to this project implements
8//! > this endpoint, and no OpenAI API key was available for testing. If you
9//! > encounter any issues, please report them on the repository.
10
11use std::path::PathBuf;
12
13use serde::Serialize;
14use url::Url;
15
16use crate::{
17    errors::OapiError,
18    images::{ImageResponseFormat, enum_to_literal},
19    rest::post::{Post, PostNoStream},
20};
21
22/// Creates a variation of a given image.
23#[derive(Debug, Serialize, Default, Clone)]
24pub struct ImageVariationRequest {
25    /// The image to use as the basis for the variation(s), as a file path.
26    ///
27    /// Must be a valid PNG file, less than 4MB, and square.
28    #[serde(skip_serializing)]
29    pub image: PathBuf,
30    /// The model to use for image generation. Only `dall-e-2` is supported
31    /// at this time.
32    #[serde(skip_serializing_if = "Option::is_none")]
33    pub model: Option<String>,
34    /// The number of images to generate. Must be between 1 and 10.
35    #[serde(skip_serializing_if = "Option::is_none")]
36    pub n: Option<u32>,
37    /// The format in which the generated images are returned. Must be one of
38    /// `url` or `b64_json`. URLs are only valid for 60 minutes after the
39    /// image has been generated.
40    #[serde(skip_serializing_if = "Option::is_none")]
41    pub response_format: Option<ImageResponseFormat>,
42    /// The size of the generated images. Must be one of `256x256`,
43    /// `512x512`, or `1024x1024`.
44    #[serde(skip_serializing_if = "Option::is_none")]
45    pub size: Option<VariationSize>,
46    /// A unique identifier representing your end-user, which can help OpenAI
47    /// to monitor and detect abuse.
48    #[serde(skip_serializing_if = "Option::is_none")]
49    pub user: Option<String>,
50}
51
52/// The size of the generated variation images.
53#[derive(Debug, Serialize, Clone, Copy)]
54pub enum VariationSize {
55    #[serde(rename = "256x256")]
56    S256x256,
57    #[serde(rename = "512x512")]
58    S512x512,
59    #[serde(rename = "1024x1024")]
60    S1024x1024,
61}
62
63impl Post for ImageVariationRequest {
64    #[inline]
65    fn is_streaming(&self) -> bool {
66        false
67    }
68
69    /// Builds the URL for the request.
70    ///
71    /// `base_url` should be like <https://api.openai.com/v1>
72    fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
73        let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
74        url.path_segments_mut()
75            .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
76            .push("images")
77            .push("variations");
78
79        Ok(url.to_string())
80    }
81}
82
83impl PostNoStream for ImageVariationRequest {
84    type Response = crate::images::ImagesResponse;
85
86    /// Sends an image variation POST request using multipart/form-data
87    /// format, following the field layout of the official SDK.
88    async fn get_response_string(
89        &self,
90        client: &reqwest::Client,
91        url: &str,
92        key: &str,
93    ) -> Result<String, OapiError> {
94        let content = tokio::fs::read(&self.image).await?;
95        let file_name = self
96            .image
97            .file_name()
98            .and_then(|name| name.to_str())
99            .ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))?
100            .to_string();
101
102        let image_part = reqwest::multipart::Part::bytes(content).file_name(file_name);
103        let mut form = reqwest::multipart::Form::new().part("image", image_part);
104
105        if let Some(model) = &self.model {
106            form = form.text("model", model.clone());
107        }
108        if let Some(n) = self.n {
109            form = form.text("n", n.to_string());
110        }
111        if let Some(response_format) = self.response_format {
112            form = form.text("response_format", enum_to_literal(&response_format)?);
113        }
114        if let Some(size) = self.size {
115            form = form.text("size", enum_to_literal(&size)?);
116        }
117        if let Some(user) = &self.user {
118            form = form.text("user", user.clone());
119        }
120
121        let response = client
122            .post(url)
123            .header("Accept", "application/json")
124            .bearer_auth(key)
125            .multipart(form)
126            .send()
127            .await?;
128
129        crate::rest::response_text_checked(response).await
130    }
131}
132
133#[cfg(test)]
134mod tests {
135    use super::*;
136
137    #[test]
138    fn test_build_url() {
139        let request = ImageVariationRequest::default();
140        let url = request.build_url("https://api.openai.com/v1/").unwrap();
141        assert_eq!(url, "https://api.openai.com/v1/images/variations");
142    }
143
144    /// Size literals serialize to their official wire values.
145    #[test]
146    fn size_literals() {
147        assert_eq!(
148            enum_to_literal(&VariationSize::S256x256).unwrap(),
149            "256x256"
150        );
151        assert_eq!(
152            enum_to_literal(&VariationSize::S1024x1024).unwrap(),
153            "1024x1024"
154        );
155    }
156}