openai-interface 0.13.0

A low-level Rust interface for the OpenAI API
Documentation
//! Creates a variation of a given image.
//!
//! Endpoint: `POST /images/variations` (multipart/form-data request / JSON
//! response). This endpoint only supports `dall-e-2`.
//!
//! > ![warn] This module is untested!
//! > No OpenAI-compatible provider accessible to this project implements
//! > this endpoint, and no OpenAI API key was available for testing. If you
//! > encounter any issues, please report them on the repository.

use std::path::PathBuf;

use serde::Serialize;
use url::Url;

use crate::{
    errors::OapiError,
    images::{ImageResponseFormat, enum_to_literal},
    rest::RequestOptions,
    rest::post::{Post, PostNoStream},
};

/// Creates a variation of a given image.
#[derive(Debug, Serialize, Default, Clone)]
pub struct ImageVariationRequest {
    /// The image to use as the basis for the variation(s), as a file path.
    ///
    /// Must be a valid PNG file, less than 4MB, and square.
    #[serde(skip_serializing)]
    pub image: PathBuf,
    /// The model to use for image generation. Only `dall-e-2` is supported
    /// at this time.
    #[serde(skip_serializing_if = "Option::is_none")]
    pub model: Option<String>,
    /// The number of images to generate. Must be between 1 and 10.
    #[serde(skip_serializing_if = "Option::is_none")]
    pub n: Option<u32>,
    /// The format in which the generated images are returned. Must be one of
    /// `url` or `b64_json`. URLs are only valid for 60 minutes after the
    /// image has been generated.
    #[serde(skip_serializing_if = "Option::is_none")]
    pub response_format: Option<ImageResponseFormat>,
    /// The size of the generated images. Must be one of `256x256`,
    /// `512x512`, or `1024x1024`.
    #[serde(skip_serializing_if = "Option::is_none")]
    pub size: Option<VariationSize>,
    /// A unique identifier representing your end-user, which can help OpenAI
    /// to monitor and detect abuse.
    #[serde(skip_serializing_if = "Option::is_none")]
    pub user: Option<String>,
    /// Additional JSON properties, sent as extra multipart text fields
    /// (strings verbatim, other JSON values serialized).
    pub extra_body_map: Option<serde_json::Map<String, serde_json::Value>>,
}

/// The size of the generated variation images.
#[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
    }

    /// Builds the URL for the request.
    ///
    /// `base_url` should be like <https://api.openai.com/v1>
    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;

    /// Sends an image variation POST request using multipart/form-data
    /// format, following the field layout of the official SDK.
    async fn get_response_string(
        &self,
        client: &reqwest::Client,
        base_url: &str,
        options: &RequestOptions,
    ) -> 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());
        }

        form = crate::rest::post::append_extra_body_map(form, &self.extra_body_map);

        let url = self.build_url(base_url)?;
        crate::rest::post::post_multipart_json(client, url, form, options).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");
    }

    /// Size literals serialize to their official wire values.
    #[test]
    fn size_literals() {
        assert_eq!(
            enum_to_literal(&VariationSize::S256x256).unwrap(),
            "256x256"
        );
        assert_eq!(
            enum_to_literal(&VariationSize::S1024x1024).unwrap(),
            "1024x1024"
        );
    }
}