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 = "pcm16")]
Pcm16,
#[serde(rename = "pcm24")]
Pcm24,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
pub enum OutputAudioFormat {
#[default]
#[serde(rename = "pcm")]
Pcm,
}
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) -> ZaiResult<String> {
if samples.len() % 2 != 0 {
return Err(RealtimeErrorKind::Protocol(
"16-bit PCM input must contain an even number of bytes".into(),
)
.into());
}
if sample_rate == 0 {
return Err(RealtimeErrorKind::Protocol("WAV sample rate must be positive".into()).into());
}
let bytes_per_sample: u32 = 2;
let channels: u32 = 1;
let byte_rate = sample_rate
.checked_mul(channels * bytes_per_sample)
.ok_or_else(|| RealtimeErrorKind::Protocol("WAV byte rate overflow".into()))?;
let block_align = u16::try_from(channels * bytes_per_sample)
.map_err(|_| RealtimeErrorKind::Protocol("WAV block alignment overflow".into()))?;
let data_len = u32::try_from(samples.len())
.map_err(|_| RealtimeErrorKind::Protocol("PCM input is too large for WAV".into()))?;
let chunk_size = data_len
.checked_add(36)
.ok_or_else(|| RealtimeErrorKind::Protocol("PCM input is too large for WAV".into()))?;
let capacity = samples
.len()
.checked_add(44)
.ok_or_else(|| RealtimeErrorKind::Protocol("PCM input is too large for WAV".into()))?;
let mut wav = Vec::with_capacity(capacity);
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);
Ok(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).unwrap();
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);
}
#[test]
fn wav_encoder_rejects_invalid_pcm_metadata() {
assert!(encode_wav_pcm_base64(&[0], 16_000).is_err());
assert!(encode_wav_pcm_base64(&[0, 0], 0).is_err());
}
#[test]
fn current_formats_use_official_wire_values() {
assert_eq!(
serde_json::to_string(&InputAudioFormat::Wav).unwrap(),
r#""wav""#
);
assert_eq!(
serde_json::to_string(&InputAudioFormat::Pcm16).unwrap(),
r#""pcm16""#
);
assert_eq!(
serde_json::to_string(&InputAudioFormat::Pcm24).unwrap(),
r#""pcm24""#
);
assert_eq!(
serde_json::to_string(&OutputAudioFormat::Pcm).unwrap(),
r#""pcm""#
);
assert!(serde_json::from_str::<InputAudioFormat>(r#""wav48""#).is_err());
assert!(serde_json::from_str::<OutputAudioFormat>(r#""mp3""#).is_err());
}
}