use serde::Serialize;
use url::Url;
use crate::{
errors::OapiError,
rest::post::{Post, PostBinary},
};
#[derive(Debug, Serialize, Default, Clone)]
pub struct SpeechRequest {
pub input: String,
pub model: String,
pub voice: SpeechVoice,
#[serde(skip_serializing_if = "Option::is_none")]
pub instructions: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<SpeechFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub speed: Option<f32>,
}
#[derive(Debug, Serialize, Clone)]
#[serde(untagged)]
pub enum SpeechVoice {
BuiltIn(String),
Custom {
id: String,
},
}
impl Default for SpeechVoice {
fn default() -> Self {
Self::BuiltIn(String::new())
}
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "snake_case")]
pub enum SpeechFormat {
Mp3,
Opus,
Aac,
Flac,
Wav,
Pcm,
}
impl Post for SpeechRequest {
#[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("speech");
Ok(url.to_string())
}
}
impl PostBinary for SpeechRequest {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn request_serialization() {
let request = SpeechRequest {
input: "The quick brown fox jumped over the lazy dog.".to_string(),
model: "gpt-4o-mini-tts".to_string(),
voice: SpeechVoice::BuiltIn("alloy".to_string()),
instructions: Some("Voice: cheerful.".to_string()),
response_format: Some(SpeechFormat::Wav),
speed: Some(1.5),
};
let json = serde_json::to_string(&request).unwrap();
assert!(
json.contains(r#""input":"The quick brown fox jumped over the lazy dog.""#),
"json: {json}"
);
assert!(
json.contains(r#""model":"gpt-4o-mini-tts""#),
"json: {json}"
);
assert!(json.contains(r#""voice":"alloy""#), "json: {json}");
assert!(
json.contains(r#""instructions":"Voice: cheerful.""#),
"json: {json}"
);
assert!(json.contains(r#""response_format":"wav""#), "json: {json}");
assert!(json.contains(r#""speed":1.5"#), "json: {json}");
}
#[test]
fn custom_voice_serialization() {
let request = SpeechRequest {
input: "Hello".to_string(),
model: "gpt-4o-mini-tts".to_string(),
voice: SpeechVoice::Custom {
id: "voice_1234".to_string(),
},
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(
json.contains(r#""voice":{"id":"voice_1234"}"#),
"json: {json}"
);
}
#[test]
fn test_build_url() {
let request = SpeechRequest::default();
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/audio/speech");
}
}