openai_interface/audio/
translations.rs1use std::path::PathBuf;
18
19use serde::{Deserialize, Serialize};
20use url::Url;
21
22use crate::{
23 audio::{AudioResponseFormat, transcriptions::TranscriptionVerbose},
24 errors::OapiError,
25 rest::post::{Post, PostNoStream},
26};
27
28#[derive(Debug, Serialize, Default, Clone)]
31pub struct TranslationRequest {
32 #[serde(skip_serializing)]
37 pub file: PathBuf,
38 pub model: String,
41 #[serde(skip_serializing_if = "Option::is_none")]
44 pub prompt: Option<String>,
45 #[serde(skip_serializing_if = "Option::is_none")]
51 pub response_format: Option<AudioResponseFormat>,
52 #[serde(skip_serializing_if = "Option::is_none")]
56 pub temperature: Option<f32>,
57}
58
59#[derive(Debug, Deserialize, Clone)]
62#[serde(untagged)]
63pub enum TranslationResponse {
64 Verbose(TranscriptionVerbose),
67 Plain(Translation),
69}
70
71#[derive(Debug, Deserialize, Clone)]
73pub struct Translation {
74 pub text: String,
76}
77
78crate::impl_from_str!(TranslationResponse);
79
80impl Post for TranslationRequest {
81 #[inline]
82 fn is_streaming(&self) -> bool {
83 false
84 }
85
86 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
90 let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
91 url.path_segments_mut()
92 .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
93 .push("audio")
94 .push("translations");
95
96 Ok(url.to_string())
97 }
98}
99
100impl PostNoStream for TranslationRequest {
101 type Response = TranslationResponse;
102
103 async fn get_response_string(
106 &self,
107 client: &reqwest::Client,
108 url: &str,
109 key: &str,
110 ) -> Result<String, OapiError> {
111 if !self.file.exists() {
112 return Err(OapiError::FileNotFoundError(self.file.clone()));
113 }
114
115 let content = tokio::fs::read(&self.file).await?;
116 let file_name = self
117 .file
118 .file_name()
119 .and_then(|name| name.to_str())
120 .ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))?
121 .to_string();
122
123 let file_part = reqwest::multipart::Part::bytes(content).file_name(file_name);
124 let mut form = reqwest::multipart::Form::new().part("file", file_part);
125
126 form = form.text("model", self.model.clone());
127
128 if let Some(prompt) = &self.prompt {
129 form = form.text("prompt", prompt.clone());
130 }
131 if let Some(response_format) = self.response_format {
132 let literal = crate::audio::enum_to_literal(&response_format)?;
133 form = form.text("response_format", literal);
134 }
135 if let Some(temperature) = self.temperature {
136 form = form.text("temperature", temperature.to_string());
137 }
138
139 let response = client
140 .post(url)
141 .header("Accept", "application/json")
142 .bearer_auth(key)
143 .multipart(form)
144 .send()
145 .await?;
146
147 crate::rest::response_text_checked(response).await
148 }
149}
150
151#[cfg(test)]
152mod tests {
153 use super::*;
154
155 #[test]
156 fn test_build_url() {
157 let request = TranslationRequest::default();
158 let url = request.build_url("https://api.openai.com/v1/").unwrap();
159 assert_eq!(url, "https://api.openai.com/v1/audio/translations");
160 }
161
162 #[test]
169 fn parse_plain_response() {
170 let content = r#"{
171 "text": "The quick brown fox jumped over the lazy dog."
172 }"#;
173
174 let response: TranslationResponse = content.parse().unwrap();
175 let TranslationResponse::Plain(translation) = response else {
176 panic!("expected plain translation");
177 };
178 assert_eq!(
179 translation.text,
180 "The quick brown fox jumped over the lazy dog."
181 );
182 }
183}