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