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`): a base64-encoded **WAV** file
6//!   (`input_audio_format: "wav"` = 16 kHz, `"wav48"` = 48 kHz).
7//! - **Output** (`response.audio.delta`): base64-encoded raw **PCM** (when
8//!   `output_audio_format: "pcm"`) or **MP3** (`"mp3"`); decoded to bytes by
9//!   the transport layer before being handed to the caller.
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/// `Wav` ⇒ 16 kHz, `Wav48` ⇒ 48 kHz (per the official protocol).
19#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
20pub enum InputAudioFormat {
21    /// 16 kHz WAV.
22    #[default]
23    #[serde(rename = "wav")]
24    Wav,
25    /// 48 kHz WAV.
26    #[serde(rename = "wav48")]
27    Wav48,
28}
29
30/// Output audio format for `session.update`.
31#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
32pub enum OutputAudioFormat {
33    /// Raw PCM (server default).
34    #[default]
35    #[serde(rename = "pcm")]
36    Pcm,
37    /// MP3 frames.
38    #[serde(rename = "mp3")]
39    Mp3,
40}
41
42/// Standard (non-URL) base64-encode, matching the wire format used by the
43/// realtime API.
44pub fn encode_base64(data: &[u8]) -> String {
45    base64::engine::general_purpose::STANDARD.encode(data)
46}
47
48/// Standard base64-decode of a realtime audio payload.
49pub 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
55/// Wrap raw 16-bit little-endian mono PCM samples in a minimal WAV header and
56/// return the base64-encoded file body for `input_audio_buffer.append`.
57///
58/// `samples` must be the raw PCM byte stream (two bytes per sample, mono).
59pub 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    // RIFF header
69    wav.extend_from_slice(b"RIFF");
70    wav.extend_from_slice(&chunk_size.to_le_bytes());
71    wav.extend_from_slice(b"WAVE");
72    // fmt chunk
73    wav.extend_from_slice(b"fmt ");
74    wav.extend_from_slice(&16u32.to_le_bytes()); // PCM fmt chunk size
75    wav.extend_from_slice(&1u16.to_le_bytes()); // PCM
76    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()); // bits per sample
81    // data chunk
82    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
89/// Base64-encode a JPEG video frame for
90/// `input_audio_buffer.append_video_frame`.
91pub 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        // 100 samples of silence = 200 bytes of PCM.
102        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()); // mono
109        assert_eq!(&wav[24..28], 16000u32.to_le_bytes()); // sample rate
110        assert_eq!(&wav[34..36], 16u16.to_le_bytes()); // bits per sample
111        assert_eq!(&wav[36..40], b"data");
112        assert_eq!(&wav[40..44], (pcm.len() as u32).to_le_bytes()); // data length
113        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}