aurum_core/remote/
openai_speech.rs1use serde::{Deserialize, Serialize};
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
10#[serde(rename_all = "snake_case")]
11pub enum SpeechResponseFormat {
12 Pcm,
14 Mp3,
16}
17
18impl SpeechResponseFormat {
19 pub fn as_str(self) -> &'static str {
20 match self {
21 Self::Pcm => "pcm",
22 Self::Mp3 => "mp3",
23 }
24 }
25}
26
27#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
29pub struct OpenAiSpeechRequest {
30 pub model: String,
31 pub input: String,
32 pub voice: String,
33 #[serde(skip_serializing_if = "Option::is_none")]
34 pub response_format: Option<SpeechResponseFormat>,
35 #[serde(skip_serializing_if = "Option::is_none")]
36 pub speed: Option<f64>,
37}
38
39impl OpenAiSpeechRequest {
40 pub fn new(
42 model: impl Into<String>,
43 input: impl Into<String>,
44 voice: impl Into<String>,
45 format: SpeechResponseFormat,
46 speed: Option<f32>,
47 ) -> Self {
48 Self {
49 model: model.into(),
50 input: input.into(),
51 voice: voice.into(),
52 response_format: Some(format),
53 speed: speed.map(|s| s as f64),
54 }
55 }
56
57 pub fn to_json_bytes(&self) -> Result<Vec<u8>, serde_json::Error> {
58 serde_json::to_vec(self)
59 }
60}
61
62pub fn parse_pcm_content_type(content_type: &str) -> Option<(u32, u16)> {
67 let ct = content_type.to_ascii_lowercase();
68 if !ct.contains("audio/pcm") && !ct.starts_with("audio/l16") {
69 return None;
70 }
71 let mut rate = None;
72 let mut channels = None;
73 for part in ct.split(';').skip(1) {
74 let part = part.trim();
75 if let Some(v) = part.strip_prefix("rate=") {
76 rate = v.trim().parse().ok();
77 } else if let Some(v) = part.strip_prefix("channels=") {
78 channels = v.trim().parse().ok();
79 }
80 }
81 let rate = rate?;
82 Some((rate, channels.unwrap_or(1)))
83}
84
85#[cfg(test)]
86mod tests {
87 use super::*;
88
89 #[test]
90 fn serializes_speech_request() {
91 let req = OpenAiSpeechRequest::new(
92 "openai/gpt-4o-mini-tts",
93 "Hello",
94 "alloy",
95 SpeechResponseFormat::Pcm,
96 Some(1.0),
97 );
98 let v: serde_json::Value = serde_json::from_slice(&req.to_json_bytes().unwrap()).unwrap();
99 assert_eq!(v["model"], "openai/gpt-4o-mini-tts");
100 assert_eq!(v["input"], "Hello");
101 assert_eq!(v["voice"], "alloy");
102 assert_eq!(v["response_format"], "pcm");
103 assert_eq!(v["speed"], 1.0);
104 let s = v.to_string();
106 assert!(!s.contains("Authorization"));
107 assert!(!s.contains("api_key"));
108 }
109
110 #[test]
111 fn omits_speed_when_none() {
112 let req = OpenAiSpeechRequest::new("m", "i", "v", SpeechResponseFormat::Mp3, None);
113 let s = serde_json::to_string(&req).unwrap();
114 assert!(!s.contains("speed"));
115 assert!(s.contains("mp3"));
116 }
117
118 #[test]
119 fn parse_pcm_content_type_rate_channels() {
120 assert_eq!(
121 parse_pcm_content_type("audio/pcm;rate=24000;channels=1"),
122 Some((24_000, 1))
123 );
124 assert_eq!(
125 parse_pcm_content_type("AUDIO/PCM; rate=16000"),
126 Some((16_000, 1))
127 );
128 assert!(parse_pcm_content_type("audio/mpeg").is_none());
129 }
130}