use std::path::PathBuf;
use serde::{Deserialize, Serialize};
use url::Url;
use crate::{
audio::{AudioResponseFormat, transcriptions::TranscriptionVerbose},
errors::OapiError,
rest::post::{Post, PostNoStream},
};
#[derive(Debug, Serialize, Default, Clone)]
pub struct TranslationRequest {
#[serde(skip_serializing)]
pub file: PathBuf,
pub model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<AudioResponseFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
}
#[derive(Debug, Deserialize, Clone)]
#[serde(untagged)]
pub enum TranslationResponse {
Verbose(TranscriptionVerbose),
Plain(Translation),
}
#[derive(Debug, Deserialize, Clone)]
pub struct Translation {
pub text: String,
}
crate::impl_from_str!(TranslationResponse);
impl Post for TranslationRequest {
#[inline]
fn is_streaming(&self) -> bool {
false
}
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("audio")
.push("translations");
Ok(url.to_string())
}
}
impl PostNoStream for TranslationRequest {
type Response = TranslationResponse;
async fn get_response_string(
&self,
client: &reqwest::Client,
url: &str,
key: &str,
) -> Result<String, OapiError> {
if !self.file.exists() {
return Err(OapiError::FileNotFoundError(self.file.clone()));
}
let content = tokio::fs::read(&self.file).await?;
let file_name = self
.file
.file_name()
.and_then(|name| name.to_str())
.ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))?
.to_string();
let file_part = reqwest::multipart::Part::bytes(content).file_name(file_name);
let mut form = reqwest::multipart::Form::new().part("file", file_part);
form = form.text("model", self.model.clone());
if let Some(prompt) = &self.prompt {
form = form.text("prompt", prompt.clone());
}
if let Some(response_format) = self.response_format {
let literal = crate::audio::enum_to_literal(&response_format)?;
form = form.text("response_format", literal);
}
if let Some(temperature) = self.temperature {
form = form.text("temperature", temperature.to_string());
}
let response = client
.post(url)
.header("Accept", "application/json")
.bearer_auth(key)
.multipart(form)
.send()
.await?;
crate::rest::response_text_checked(response).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_build_url() {
let request = TranslationRequest::default();
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/audio/translations");
}
#[test]
fn parse_plain_response() {
let content = r#"{
"text": "The quick brown fox jumped over the lazy dog."
}"#;
let response: TranslationResponse = content.parse().unwrap();
let TranslationResponse::Plain(translation) = response else {
panic!("expected plain translation");
};
assert_eq!(
translation.text,
"The quick brown fox jumped over the lazy dog."
);
}
}