1use std::io::Read;
4
5use super::{
6 binary, Wav, WavEncodingCode, WavError, WavType, WavWaveMetadata, WavWrapperKind,
7 DATA_CHUNK_ID, FMT_CHUNK_ID, MP3_IN_WAV_HEADER_SIZE, MP3_IN_WAV_RIFF_SIZE, RIFF_MAGIC,
8 SFX_HEADER_SIZE, VO_HEADER_SIZE, WAVE_MAGIC,
9};
10
11#[cfg_attr(
18 feature = "tracing",
19 tracing::instrument(level = "debug", skip(reader))
20)]
21pub fn read_wav<R: Read>(reader: &mut R) -> Result<Wav, WavError> {
22 let mut bytes = Vec::new();
23 reader.read_to_end(&mut bytes)?;
24 read_wav_from_bytes(&bytes)
25}
26
27#[cfg_attr(
38 feature = "tracing",
39 tracing::instrument(level = "debug", skip(bytes), fields(bytes_len = bytes.len()))
40)]
41pub fn read_wav_from_bytes(bytes: &[u8]) -> Result<Wav, WavError> {
42 let (kind, skip_size) = detect_wrapper(bytes);
43 if skip_size > bytes.len() {
44 return Err(WavError::InvalidHeader(format!(
45 "wrapper skip offset {skip_size} exceeds file length {}",
46 bytes.len()
47 )));
48 }
49
50 let payload = &bytes[skip_size..];
51 let wav_type = match kind {
52 WavWrapperKind::SfxHeader => WavType::Sfx,
53 WavWrapperKind::Standard => WavType::Standard,
54 _ => WavType::Vo, };
56
57 if kind == WavWrapperKind::Mp3InWav {
58 return Ok(Wav::new_mp3(wav_type, payload.to_vec()));
59 }
60
61 parse_wave_payload(payload, wav_type)
62}
63
64fn detect_wrapper(bytes: &[u8]) -> (WavWrapperKind, usize) {
65 if bytes.len() < 12 {
66 return (WavWrapperKind::Standard, 0);
67 }
68
69 if bytes.starts_with(&RIFF_MAGIC) {
70 if bytes.len() >= VO_HEADER_SIZE + 4
71 && bytes[VO_HEADER_SIZE..VO_HEADER_SIZE + 4] == RIFF_MAGIC
72 {
73 return (WavWrapperKind::VoHeader, VO_HEADER_SIZE);
74 }
75
76 if let Ok(riff_size) = binary::read_u32(bytes, 4) {
77 if riff_size == MP3_IN_WAV_RIFF_SIZE {
78 return (WavWrapperKind::Mp3InWav, MP3_IN_WAV_HEADER_SIZE);
79 }
80 }
81
82 return (WavWrapperKind::Standard, 0);
83 }
84
85 if bytes.len() >= SFX_HEADER_SIZE + 4
93 && bytes[SFX_HEADER_SIZE..SFX_HEADER_SIZE + 4] == RIFF_MAGIC
94 {
95 return (WavWrapperKind::SfxHeader, SFX_HEADER_SIZE);
96 }
97
98 (WavWrapperKind::Standard, 0)
99}
100
101fn parse_wave_payload(bytes: &[u8], wav_type: WavType) -> Result<Wav, WavError> {
102 if bytes.len() < 12 {
103 return Err(WavError::InvalidHeader(
104 "WAVE payload shorter than 12-byte RIFF header".into(),
105 ));
106 }
107 if !bytes.starts_with(&RIFF_MAGIC) {
108 return Err(WavError::InvalidHeader(
109 "missing RIFF magic after wrapper normalization".into(),
110 ));
111 }
112 if bytes[8..12] != WAVE_MAGIC {
113 return Err(WavError::InvalidHeader(
114 "missing WAVE magic in RIFF header".into(),
115 ));
116 }
117
118 let mut cursor = 12usize;
119 let mut fmt_fields: Option<(WavEncodingCode, u16, u32, u32, u16, u16)> = None;
120 let mut data_chunk: Option<Vec<u8>> = None;
121
122 while cursor + 8 <= bytes.len() {
123 let chunk_id = [
124 bytes[cursor],
125 bytes[cursor + 1],
126 bytes[cursor + 2],
127 bytes[cursor + 3],
128 ];
129 let chunk_size_u32 = binary::read_u32(bytes, cursor + 4).map_err(|_| {
130 WavError::InvalidChunk(format!(
131 "unable to read chunk size at offset {}",
132 cursor + 4
133 ))
134 })?;
135 let chunk_size = usize::try_from(chunk_size_u32).map_err(|_| {
136 WavError::InvalidChunk(format!("chunk size {chunk_size_u32} does not fit in usize"))
137 })?;
138 cursor += 8;
139
140 let chunk_end = cursor.checked_add(chunk_size).ok_or_else(|| {
141 WavError::InvalidChunk(format!("chunk size overflow at chunk offset {cursor}"))
142 })?;
143 if chunk_end > bytes.len() {
144 return Err(WavError::InvalidChunk(format!(
145 "chunk at offset {} exceeds file bounds",
146 cursor - 8
147 )));
148 }
149 let chunk_data = &bytes[cursor..chunk_end];
150
151 if chunk_id == FMT_CHUNK_ID {
152 if chunk_data.len() < 16 {
153 return Err(WavError::InvalidChunk(
154 "`fmt ` chunk shorter than 16 bytes".into(),
155 ));
156 }
157
158 let encoding =
159 WavEncodingCode::from_raw(u16::from_le_bytes([chunk_data[0], chunk_data[1]]));
160 let channels = u16::from_le_bytes([chunk_data[2], chunk_data[3]]);
161 let sample_rate =
162 u32::from_le_bytes([chunk_data[4], chunk_data[5], chunk_data[6], chunk_data[7]]);
163 let bytes_per_sec =
164 u32::from_le_bytes([chunk_data[8], chunk_data[9], chunk_data[10], chunk_data[11]]);
165 let block_align = u16::from_le_bytes([chunk_data[12], chunk_data[13]]);
166 let bits_per_sample = u16::from_le_bytes([chunk_data[14], chunk_data[15]]);
167 fmt_fields = Some((
168 encoding,
169 channels,
170 sample_rate,
171 bytes_per_sec,
172 block_align,
173 bits_per_sample,
174 ));
175 } else if chunk_id == DATA_CHUNK_ID {
176 data_chunk = Some(chunk_data.to_vec());
177 break;
178 }
179
180 cursor = chunk_end;
181 if chunk_size % 2 == 1 && cursor < bytes.len() {
182 cursor += 1;
183 }
184 }
185
186 let (encoding, channels, sample_rate, bytes_per_sec, block_align, bits_per_sample) =
187 fmt_fields.ok_or_else(|| WavError::InvalidChunk("missing required `fmt ` chunk".into()))?;
188 let data =
189 data_chunk.ok_or_else(|| WavError::InvalidChunk("missing required `data` chunk".into()))?;
190
191 Ok(Wav::new_wave(
192 wav_type,
193 WavWaveMetadata {
194 encoding,
195 channels,
196 sample_rate,
197 bytes_per_sec,
198 block_align,
199 bits_per_sample,
200 },
201 data,
202 ))
203}
204
205#[cfg(test)]
206mod tests {
207 use super::*;
208 use crate::binary::{DecodeBinary, EncodeBinary};
209 use crate::wav::{
210 write_wav_to_vec, write_wav_to_vec_with_options, WavAudioFormat, WavEncoding, WavWriteMode,
211 WavWriteOptions, SFX_MAGIC,
212 };
213
214 fn sample_wave() -> Wav {
215 Wav::new_wave(
216 WavType::Vo,
217 WavWaveMetadata {
218 encoding: WavEncodingCode::from(WavEncoding::Pcm),
219 channels: 1,
220 sample_rate: 22_050,
221 bytes_per_sec: 44_100,
222 block_align: 2,
223 bits_per_sample: 16,
224 },
225 vec![1, 2, 3, 4, 5, 6],
226 )
227 }
228
229 #[test]
230 fn roundtrip_clean_wave_payload() {
231 let wav = sample_wave();
232 let bytes = write_wav_to_vec_with_options(
233 &wav,
234 WavWriteOptions {
235 mode: WavWriteMode::Clean,
236 },
237 )
238 .expect("write clean WAV");
239 let parsed = read_wav_from_bytes(&bytes).expect("read clean WAV");
240
241 assert_eq!(parsed.audio_format, WavAudioFormat::Wave);
242 assert_eq!(parsed.encoding.known(), Some(WavEncoding::Pcm));
243 assert_eq!(parsed.channels, 1);
244 assert_eq!(parsed.sample_rate, 22_050);
245 assert_eq!(parsed.data, wav.data);
246 assert_eq!(parsed.wav_type, WavType::Standard);
248 }
249
250 #[test]
251 fn writer_is_deterministic_for_synthetic_wave() {
252 let wav = sample_wave();
253 let first = write_wav_to_vec(&wav).expect("first write should succeed");
254 let second = write_wav_to_vec(&wav).expect("second write should succeed");
255 assert_eq!(first, second);
256 }
257
258 #[test]
259 fn roundtrip_game_sfx_wrapper() {
260 let mut wav = sample_wave();
261 wav.wav_type = WavType::Sfx;
262
263 let bytes = write_wav_to_vec(&wav).expect("write game WAV");
264 assert!(bytes.starts_with(&SFX_MAGIC));
265
266 let parsed = read_wav_from_bytes(&bytes).expect("read SFX WAV");
267 assert_eq!(parsed.wav_type, WavType::Sfx);
268 assert_eq!(parsed.audio_format, WavAudioFormat::Wave);
269 assert_eq!(parsed.data, wav.data);
270 }
271
272 #[test]
273 fn wrapped_sfx_is_detected_without_the_vanilla_prefix() {
274 let mut wav = sample_wave();
279 wav.wav_type = WavType::Sfx;
280 let mut bytes = write_wav_to_vec(&wav).expect("write game WAV");
281 bytes[0..4].copy_from_slice(b"\x00\x11\x22\x33");
282 assert!(!bytes.starts_with(&SFX_MAGIC));
283
284 let parsed = read_wav_from_bytes(&bytes).expect("read SFX WAV with a foreign prefix");
285 assert_eq!(parsed.wav_type, WavType::Sfx);
286 assert_eq!(parsed.data, wav.data);
287 }
288
289 #[test]
290 fn a_file_with_riff_nowhere_is_not_treated_as_wrapped() {
291 let bytes = vec![0u8; SFX_HEADER_SIZE + 64];
294 let err = read_wav_from_bytes(&bytes);
295 assert!(err.is_err(), "a file with no RIFF anywhere must not parse");
296 }
297
298 #[test]
299 fn roundtrip_game_vo_wrapper() {
300 let wav = sample_wave();
301 let bytes = write_wav_to_vec(&wav).expect("write game WAV");
302 assert!(bytes.starts_with(&RIFF_MAGIC));
303 assert_eq!(&bytes[VO_HEADER_SIZE..VO_HEADER_SIZE + 4], &RIFF_MAGIC);
304
305 let parsed = read_wav_from_bytes(&bytes).expect("read VO WAV");
306 assert_eq!(parsed.wav_type, WavType::Vo);
307 assert_eq!(parsed.audio_format, WavAudioFormat::Wave);
308 assert_eq!(parsed.data, wav.data);
309 }
310
311 #[test]
312 fn detects_and_unwraps_mp3_in_wav() {
313 let mp3_payload = vec![0x49, 0x44, 0x33, 0x04, 0x00];
314 let mut wrapped = vec![0_u8; MP3_IN_WAV_HEADER_SIZE];
315 wrapped[0..4].copy_from_slice(&RIFF_MAGIC);
316 wrapped[4..8].copy_from_slice(&MP3_IN_WAV_RIFF_SIZE.to_le_bytes());
317 wrapped[8..12].copy_from_slice(&WAVE_MAGIC);
318 wrapped.extend_from_slice(&mp3_payload);
319
320 let parsed = read_wav_from_bytes(&wrapped).expect("read MP3-in-WAV");
321 assert_eq!(parsed.wav_type, WavType::Vo);
322 assert_eq!(parsed.audio_format, WavAudioFormat::Mp3);
323 assert_eq!(parsed.encoding.known(), Some(WavEncoding::Mp3));
324 assert_eq!(parsed.data, mp3_payload);
325 }
326
327 #[test]
328 fn clean_mode_mp3_writes_raw_payload() {
329 let wav = Wav::new_mp3(WavType::Vo, vec![0x01, 0x02, 0x03]);
330 let bytes = write_wav_to_vec_with_options(
331 &wav,
332 WavWriteOptions {
333 mode: WavWriteMode::Clean,
334 },
335 )
336 .expect("write clean MP3");
337 assert_eq!(bytes, vec![0x01, 0x02, 0x03]);
338 }
339
340 #[test]
341 fn rejects_non_riff_wave_payload() {
342 let err = read_wav_from_bytes(b"not-riff").expect_err("must fail");
343 assert!(matches!(err, WavError::InvalidHeader(_)));
344 }
345
346 #[test]
347 fn rejects_truncated_header() {
348 let err = read_wav_from_bytes(&RIFF_MAGIC).expect_err("must fail");
349 assert!(matches!(err, WavError::InvalidHeader(_)));
350 }
351
352 #[test]
353 fn rejects_missing_fmt_chunk() {
354 let mut bytes = Vec::new();
355 bytes.extend_from_slice(&RIFF_MAGIC);
356 bytes.extend_from_slice(&(12_u32 + 8 + 4).to_le_bytes());
357 bytes.extend_from_slice(&WAVE_MAGIC);
358 bytes.extend_from_slice(&DATA_CHUNK_ID);
359 bytes.extend_from_slice(&(4_u32).to_le_bytes());
360 bytes.extend_from_slice(&[0, 1, 2, 3]);
361
362 let err = read_wav_from_bytes(&bytes).expect_err("must fail");
363 assert!(matches!(err, WavError::InvalidChunk(_)));
364 }
365
366 #[test]
367 fn rejects_truncated_chunk_payload() {
368 let mut bytes = Vec::new();
369 bytes.extend_from_slice(&RIFF_MAGIC);
370 bytes.extend_from_slice(&(12_u32 + 8 + 16).to_le_bytes());
371 bytes.extend_from_slice(&WAVE_MAGIC);
372 bytes.extend_from_slice(&FMT_CHUNK_ID);
373 bytes.extend_from_slice(&(16_u32).to_le_bytes());
374 bytes.extend_from_slice(&[1, 0, 1, 0]); let err = read_wav_from_bytes(&bytes).expect_err("must fail");
377 assert!(matches!(err, WavError::InvalidChunk(_)));
378 }
379
380 #[test]
381 fn decode_encode_traits_roundtrip() {
382 let mut wav = sample_wave();
383 wav.wav_type = WavType::Sfx;
384
385 let bytes = wav.encode_binary().expect("encode");
386 let decoded = Wav::decode_binary(&bytes).expect("decode");
387
388 assert_eq!(decoded.wav_type, WavType::Sfx);
389 assert_eq!(decoded.audio_format, WavAudioFormat::Wave);
390 assert_eq!(decoded.data, wav.data);
391 }
392}