openai_interface/images/
variation.rs1use 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#[derive(Debug, Serialize, Default, Clone)]
25pub struct ImageVariationRequest {
26 #[serde(skip_serializing)]
30 pub image: PathBuf,
31 #[serde(skip_serializing_if = "Option::is_none")]
34 pub model: Option<String>,
35 #[serde(skip_serializing_if = "Option::is_none")]
37 pub n: Option<u32>,
38 #[serde(skip_serializing_if = "Option::is_none")]
42 pub response_format: Option<ImageResponseFormat>,
43 #[serde(skip_serializing_if = "Option::is_none")]
46 pub size: Option<VariationSize>,
47 #[serde(skip_serializing_if = "Option::is_none")]
50 pub user: Option<String>,
51}
52
53#[derive(Debug, Serialize, Clone, Copy)]
55pub enum VariationSize {
56 #[serde(rename = "256x256")]
57 S256x256,
58 #[serde(rename = "512x512")]
59 S512x512,
60 #[serde(rename = "1024x1024")]
61 S1024x1024,
62}
63
64impl Post for ImageVariationRequest {
65 #[inline]
66 fn is_streaming(&self) -> bool {
67 false
68 }
69
70 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
74 let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
75 url.path_segments_mut()
76 .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
77 .push("images")
78 .push("variations");
79
80 Ok(url.to_string())
81 }
82}
83
84impl PostNoStream for ImageVariationRequest {
85 type Response = crate::images::ImagesResponse;
86
87 async fn get_response_string(
90 &self,
91 client: &reqwest::Client,
92 base_url: &str,
93 options: &RequestOptions,
94 ) -> Result<String, OapiError> {
95 let content = tokio::fs::read(&self.image).await?;
96 let file_name = self
97 .image
98 .file_name()
99 .and_then(|name| name.to_str())
100 .ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))?
101 .to_string();
102
103 let image_part = reqwest::multipart::Part::bytes(content).file_name(file_name);
104 let mut form = reqwest::multipart::Form::new().part("image", image_part);
105
106 if let Some(model) = &self.model {
107 form = form.text("model", model.clone());
108 }
109 if let Some(n) = self.n {
110 form = form.text("n", n.to_string());
111 }
112 if let Some(response_format) = self.response_format {
113 form = form.text("response_format", enum_to_literal(&response_format)?);
114 }
115 if let Some(size) = self.size {
116 form = form.text("size", enum_to_literal(&size)?);
117 }
118 if let Some(user) = &self.user {
119 form = form.text("user", user.clone());
120 }
121
122 let url = self.build_url(base_url)?;
123 crate::rest::post::post_multipart_json(client, url, form, options).await
124 }
125}
126
127#[cfg(test)]
128mod tests {
129 use super::*;
130
131 #[test]
132 fn test_build_url() {
133 let request = ImageVariationRequest::default();
134 let url = request.build_url("https://api.openai.com/v1/").unwrap();
135 assert_eq!(url, "https://api.openai.com/v1/images/variations");
136 }
137
138 #[test]
140 fn size_literals() {
141 assert_eq!(
142 enum_to_literal(&VariationSize::S256x256).unwrap(),
143 "256x256"
144 );
145 assert_eq!(
146 enum_to_literal(&VariationSize::S1024x1024).unwrap(),
147 "1024x1024"
148 );
149 }
150}