Skip to main content

zai_rs/realtime/
audio.rs

1//! Audio format + base64 helpers for the realtime API.
2//!
3//! The GLM-Realtime protocol exchanges audio as base64-encoded payloads:
4//!
5//! - **Input** (`input_audio_buffer.append`): base64-encoded WAV or raw PCM.
6//!   Raw PCM declares its sample rate in `input_audio_format` (`"pcm16"` for
7//!   16 kHz or `"pcm24"` for 24 kHz).
8//! - **Output** (`response.audio.delta`): base64-encoded raw 24 kHz, mono,
9//!   16-bit PCM. The session decodes each delta before handing it to callers.
10
11use base64::Engine;
12use serde::{Deserialize, Serialize};
13
14use crate::{ZaiResult, client::error::RealtimeErrorKind};
15
16/// Input audio format for `session.update`.
17///
18/// The current GLM-Realtime protocol accepts WAV or raw PCM at 16/24 kHz.
19#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
20pub enum InputAudioFormat {
21    /// A WAV container. [`RealtimeSession::send_audio`](super::RealtimeSession::send_audio)
22    /// creates a mono, 16-bit, 16 kHz WAV when this format is selected.
23    #[default]
24    #[serde(rename = "wav")]
25    Wav,
26    /// Raw 16-bit little-endian mono PCM sampled at 16 kHz.
27    #[serde(rename = "pcm16")]
28    Pcm16,
29    /// Raw 16-bit little-endian mono PCM sampled at 24 kHz.
30    #[serde(rename = "pcm24")]
31    Pcm24,
32}
33
34/// Output audio format for `session.update`.
35#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
36pub enum OutputAudioFormat {
37    /// Raw 24 kHz, mono, 16-bit PCM (the only current output format).
38    #[default]
39    #[serde(rename = "pcm")]
40    Pcm,
41}
42
43/// Standard (non-URL) base64-encode, matching the wire format used by the
44/// realtime API.
45pub fn encode_base64(data: &[u8]) -> String {
46    base64::engine::general_purpose::STANDARD.encode(data)
47}
48
49/// Standard base64-decode of a realtime audio payload.
50pub 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
56/// Wrap raw 16-bit little-endian mono PCM samples in a minimal WAV header and
57/// return the base64-encoded file body for `input_audio_buffer.append`.
58///
59/// `samples` must be the raw PCM byte stream (two bytes per sample, mono).
60pub 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    // RIFF header
90    wav.extend_from_slice(b"RIFF");
91    wav.extend_from_slice(&chunk_size.to_le_bytes());
92    wav.extend_from_slice(b"WAVE");
93    // fmt chunk
94    wav.extend_from_slice(b"fmt ");
95    wav.extend_from_slice(&16u32.to_le_bytes()); // PCM fmt chunk size
96    wav.extend_from_slice(&1u16.to_le_bytes()); // PCM
97    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()); // bits per sample
102    // data chunk
103    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
110/// Base64-encode a JPEG video frame for
111/// `input_audio_buffer.append_video_frame`.
112pub 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        // 100 samples of silence = 200 bytes of PCM.
123        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()); // mono
130        assert_eq!(&wav[24..28], 16000u32.to_le_bytes()); // sample rate
131        assert_eq!(&wav[34..36], 16u16.to_le_bytes()); // bits per sample
132        assert_eq!(&wav[36..40], b"data");
133        assert_eq!(&wav[40..44], (pcm.len() as u32).to_le_bytes()); // data length
134        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}