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 pub extra_body_map: Option<serde_json::Map<String, serde_json::Value>>,
54}
55
56#[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 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 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 #[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}