use base64::Engine;
use serde::{Deserialize, Serialize};
use crate::{ZaiResult, client::error::RealtimeErrorKind};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
pub enum InputAudioFormat {
#[default]
#[serde(rename = "wav")]
Wav,
#[serde(rename = "wav48")]
Wav48,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
pub enum OutputAudioFormat {
#[default]
#[serde(rename = "pcm")]
Pcm,
#[serde(rename = "mp3")]
Mp3,
}
pub fn encode_base64(data: &[u8]) -> String {
base64::engine::general_purpose::STANDARD.encode(data)
}
pub fn decode_base64(s: &str) -> ZaiResult<Vec<u8>> {
base64::engine::general_purpose::STANDARD
.decode(s)
.map_err(|e| RealtimeErrorKind::Protocol(format!("base64 decode failed: {e}")).into())
}
pub fn encode_wav_pcm_base64(samples: &[u8], sample_rate: u32) -> String {
let bytes_per_sample: u32 = 2;
let channels: u32 = 1;
let byte_rate = sample_rate * channels * bytes_per_sample;
let block_align = (channels * bytes_per_sample) as u16;
let data_len = samples.len() as u32;
let chunk_size = 36 + data_len;
let mut wav = Vec::with_capacity(44 + samples.len());
wav.extend_from_slice(b"RIFF");
wav.extend_from_slice(&chunk_size.to_le_bytes());
wav.extend_from_slice(b"WAVE");
wav.extend_from_slice(b"fmt ");
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());
wav.extend_from_slice(&sample_rate.to_le_bytes());
wav.extend_from_slice(&byte_rate.to_le_bytes());
wav.extend_from_slice(&block_align.to_le_bytes());
wav.extend_from_slice(&((bytes_per_sample * 8) as u16).to_le_bytes()); wav.extend_from_slice(b"data");
wav.extend_from_slice(&data_len.to_le_bytes());
wav.extend_from_slice(samples);
encode_base64(&wav)
}
pub fn encode_jpeg_frame_base64(jpg: &[u8]) -> String {
encode_base64(jpg)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn wav_round_trip_has_valid_header() {
let pcm = vec![0u8; 200];
let wav_b64 = encode_wav_pcm_base64(&pcm, 16000);
let wav = decode_base64(&wav_b64).unwrap();
assert_eq!(&wav[0..4], b"RIFF");
assert_eq!(&wav[8..12], b"WAVE");
assert_eq!(&wav[12..16], b"fmt ");
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");
assert_eq!(&wav[40..44], (pcm.len() as u32).to_le_bytes()); assert_eq!(&wav[44..], &pcm[..]);
}
#[test]
fn base64_round_trip() {
let data = b"hello realtime";
assert_eq!(decode_base64(&encode_base64(data)).unwrap(), data);
}
}