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]
23 #[serde(rename = "wav")]
24 Wav,
25 #[serde(rename = "wav48")]
27 Wav48,
28}
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
32pub enum OutputAudioFormat {
33 #[default]
35 #[serde(rename = "pcm")]
36 Pcm,
37 #[serde(rename = "mp3")]
39 Mp3,
40}
41
42pub fn encode_base64(data: &[u8]) -> String {
45 base64::engine::general_purpose::STANDARD.encode(data)
46}
47
48pub fn decode_base64(s: &str) -> ZaiResult<Vec<u8>> {
50 base64::engine::general_purpose::STANDARD
51 .decode(s)
52 .map_err(|e| RealtimeErrorKind::Protocol(format!("base64 decode failed: {e}")).into())
53}
54
55pub fn encode_wav_pcm_base64(samples: &[u8], sample_rate: u32) -> String {
60 let bytes_per_sample: u32 = 2;
61 let channels: u32 = 1;
62 let byte_rate = sample_rate * channels * bytes_per_sample;
63 let block_align = (channels * bytes_per_sample) as u16;
64 let data_len = samples.len() as u32;
65 let chunk_size = 36 + data_len;
66
67 let mut wav = Vec::with_capacity(44 + samples.len());
68 wav.extend_from_slice(b"RIFF");
70 wav.extend_from_slice(&chunk_size.to_le_bytes());
71 wav.extend_from_slice(b"WAVE");
72 wav.extend_from_slice(b"fmt ");
74 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());
77 wav.extend_from_slice(&sample_rate.to_le_bytes());
78 wav.extend_from_slice(&byte_rate.to_le_bytes());
79 wav.extend_from_slice(&block_align.to_le_bytes());
80 wav.extend_from_slice(&((bytes_per_sample * 8) as u16).to_le_bytes()); wav.extend_from_slice(b"data");
83 wav.extend_from_slice(&data_len.to_le_bytes());
84 wav.extend_from_slice(samples);
85
86 encode_base64(&wav)
87}
88
89pub fn encode_jpeg_frame_base64(jpg: &[u8]) -> String {
92 encode_base64(jpg)
93}
94
95#[cfg(test)]
96mod tests {
97 use super::*;
98
99 #[test]
100 fn wav_round_trip_has_valid_header() {
101 let pcm = vec![0u8; 200];
103 let wav_b64 = encode_wav_pcm_base64(&pcm, 16000);
104 let wav = decode_base64(&wav_b64).unwrap();
105 assert_eq!(&wav[0..4], b"RIFF");
106 assert_eq!(&wav[8..12], b"WAVE");
107 assert_eq!(&wav[12..16], b"fmt ");
108 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");
112 assert_eq!(&wav[40..44], (pcm.len() as u32).to_le_bytes()); assert_eq!(&wav[44..], &pcm[..]);
114 }
115
116 #[test]
117 fn base64_round_trip() {
118 let data = b"hello realtime";
119 assert_eq!(decode_base64(&encode_base64(data)).unwrap(), data);
120 }
121}