1use base64::Engine;
12use serde::{Deserialize, Serialize};
13
14use crate::{ZaiResult, client::error::RealtimeErrorKind};
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
20pub enum InputAudioFormat {
21 #[default]
24 #[serde(rename = "wav")]
25 Wav,
26 #[serde(rename = "pcm16")]
28 Pcm16,
29 #[serde(rename = "pcm24")]
31 Pcm24,
32}
33
34#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
36pub enum OutputAudioFormat {
37 #[default]
39 #[serde(rename = "pcm")]
40 Pcm,
41}
42
43pub fn encode_base64(data: &[u8]) -> String {
46 base64::engine::general_purpose::STANDARD.encode(data)
47}
48
49pub fn decode_base64(s: &str) -> ZaiResult<Vec<u8>> {
51 base64::engine::general_purpose::STANDARD
52 .decode(s)
53 .map_err(|e| RealtimeErrorKind::Protocol(format!("base64 decode failed: {e}")).into())
54}
55
56pub fn encode_wav_pcm_base64(samples: &[u8], sample_rate: u32) -> ZaiResult<String> {
61 if samples.len() % 2 != 0 {
62 return Err(RealtimeErrorKind::Protocol(
63 "16-bit PCM input must contain an even number of bytes".into(),
64 )
65 .into());
66 }
67 if sample_rate == 0 {
68 return Err(RealtimeErrorKind::Protocol("WAV sample rate must be positive".into()).into());
69 }
70
71 let bytes_per_sample: u32 = 2;
72 let channels: u32 = 1;
73 let byte_rate = sample_rate
74 .checked_mul(channels * bytes_per_sample)
75 .ok_or_else(|| RealtimeErrorKind::Protocol("WAV byte rate overflow".into()))?;
76 let block_align = u16::try_from(channels * bytes_per_sample)
77 .map_err(|_| RealtimeErrorKind::Protocol("WAV block alignment overflow".into()))?;
78 let data_len = u32::try_from(samples.len())
79 .map_err(|_| RealtimeErrorKind::Protocol("PCM input is too large for WAV".into()))?;
80 let chunk_size = data_len
81 .checked_add(36)
82 .ok_or_else(|| RealtimeErrorKind::Protocol("PCM input is too large for WAV".into()))?;
83 let capacity = samples
84 .len()
85 .checked_add(44)
86 .ok_or_else(|| RealtimeErrorKind::Protocol("PCM input is too large for WAV".into()))?;
87
88 let mut wav = Vec::with_capacity(capacity);
89 wav.extend_from_slice(b"RIFF");
91 wav.extend_from_slice(&chunk_size.to_le_bytes());
92 wav.extend_from_slice(b"WAVE");
93 wav.extend_from_slice(b"fmt ");
95 wav.extend_from_slice(&16u32.to_le_bytes()); wav.extend_from_slice(&1u16.to_le_bytes()); wav.extend_from_slice(&(channels as u16).to_le_bytes());
98 wav.extend_from_slice(&sample_rate.to_le_bytes());
99 wav.extend_from_slice(&byte_rate.to_le_bytes());
100 wav.extend_from_slice(&block_align.to_le_bytes());
101 wav.extend_from_slice(&((bytes_per_sample * 8) as u16).to_le_bytes()); wav.extend_from_slice(b"data");
104 wav.extend_from_slice(&data_len.to_le_bytes());
105 wav.extend_from_slice(samples);
106
107 Ok(encode_base64(&wav))
108}
109
110pub fn encode_jpeg_frame_base64(jpg: &[u8]) -> String {
113 encode_base64(jpg)
114}
115
116#[cfg(test)]
117mod tests {
118 use super::*;
119
120 #[test]
121 fn wav_round_trip_has_valid_header() {
122 let pcm = vec![0u8; 200];
124 let wav_b64 = encode_wav_pcm_base64(&pcm, 16000).unwrap();
125 let wav = decode_base64(&wav_b64).unwrap();
126 assert_eq!(&wav[0..4], b"RIFF");
127 assert_eq!(&wav[8..12], b"WAVE");
128 assert_eq!(&wav[12..16], b"fmt ");
129 assert_eq!(&wav[22..24], 1u16.to_le_bytes()); assert_eq!(&wav[24..28], 16000u32.to_le_bytes()); assert_eq!(&wav[34..36], 16u16.to_le_bytes()); assert_eq!(&wav[36..40], b"data");
133 assert_eq!(&wav[40..44], (pcm.len() as u32).to_le_bytes()); assert_eq!(&wav[44..], &pcm[..]);
135 }
136
137 #[test]
138 fn base64_round_trip() {
139 let data = b"hello realtime";
140 assert_eq!(decode_base64(&encode_base64(data)).unwrap(), data);
141 }
142
143 #[test]
144 fn wav_encoder_rejects_invalid_pcm_metadata() {
145 assert!(encode_wav_pcm_base64(&[0], 16_000).is_err());
146 assert!(encode_wav_pcm_base64(&[0, 0], 0).is_err());
147 }
148
149 #[test]
150 fn current_formats_use_official_wire_values() {
151 assert_eq!(
152 serde_json::to_string(&InputAudioFormat::Wav).unwrap(),
153 r#""wav""#
154 );
155 assert_eq!(
156 serde_json::to_string(&InputAudioFormat::Pcm16).unwrap(),
157 r#""pcm16""#
158 );
159 assert_eq!(
160 serde_json::to_string(&InputAudioFormat::Pcm24).unwrap(),
161 r#""pcm24""#
162 );
163 assert_eq!(
164 serde_json::to_string(&OutputAudioFormat::Pcm).unwrap(),
165 r#""pcm""#
166 );
167 assert!(serde_json::from_str::<InputAudioFormat>(r#""wav48""#).is_err());
168 assert!(serde_json::from_str::<OutputAudioFormat>(r#""mp3""#).is_err());
169 }
170}