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::post::{Post, PostNoStream},
20};
21
22#[derive(Debug, Serialize, Default, Clone)]
24pub struct ImageVariationRequest {
25 #[serde(skip_serializing)]
29 pub image: PathBuf,
30 #[serde(skip_serializing_if = "Option::is_none")]
33 pub model: Option<String>,
34 #[serde(skip_serializing_if = "Option::is_none")]
36 pub n: Option<u32>,
37 #[serde(skip_serializing_if = "Option::is_none")]
41 pub response_format: Option<ImageResponseFormat>,
42 #[serde(skip_serializing_if = "Option::is_none")]
45 pub size: Option<VariationSize>,
46 #[serde(skip_serializing_if = "Option::is_none")]
49 pub user: Option<String>,
50}
51
52#[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 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 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 #[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}