lasprs 0.14.1

Library for Acoustic Signal Processing (Rust edition, with optional Python bindings via pyo3)
//! WAV file export: write LASP measurements to .wav files.
use super::wav_common::*;
use super::{Result, *};
use snafu::prelude::*;

/// Desired output sample type for WAV export.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u32)]
#[cfg_attr(
    feature = "python-bindings",
    gen_stub_pyclass_enum,
    pyclass(eq, eq_int, from_py_object)
)]
pub enum ExportWavDtype {
    /// 32-bit float PCM
    F32,
    /// Signed 16-bit PCM
    I16,
    /// Signed 32-bit PCM
    I32,
}

#[cfg(feature = "python-bindings")]
#[pymethods]
impl ExportWavDtype {
    fn __str__(&self) -> &'static str {
        match self {
            ExportWavDtype::F32 => "32-bit float PCM",
            ExportWavDtype::I16 => "16-bit int PCM",
            ExportWavDtype::I32 => "32-bit int PCM",
        }
    }
}

/// Settings for exporting a measurement to WAV.
///
/// By default the exported WAV inherits the data type of the stored audio.
#[derive(Debug, Clone)]
#[cfg_attr(feature = "python-bindings", gen_stub_pyclass, pyclass(from_py_object))]
pub struct ExportWavSettings {
    /// Desired output sample type. When `None` the stored type is kept.
    pub dtype: Option<ExportWavDtype>,
    /// If `true`, normalise the signal so that `max(|audio|) == 1.0`.
    /// Normalisation is performed **before** any type conversion.
    pub normalize: bool,
    /// If `true`, overwrite the output `.wav` file when it already exists.
    /// Otherwise an error is returned on collision.
    pub overwrite: bool,
    /// If `true`, allow lossy conversion (e.g. float → int) without warning.
    /// When `false` such conversions produce an error.
    pub allow_lossy: bool,
}

#[cfg(feature = "python-bindings")]
#[cfg_attr(feature = "python-bindings", gen_stub_pymethods, pymethods)]
impl ExportWavSettings {
    #[new]
    fn new_py(
        dtype: Option<ExportWavDtype>,
        normalize: bool,
        overwrite: bool,
        allow_lossy: bool,
    ) -> Self {
        Self {
            dtype,
            normalize,
            overwrite,
            allow_lossy,
        }
    }
}

fn dtype_rank(dt: &DataType) -> u8 {
    match dt {
        DataType::I8 => 0,
        DataType::I16 => 1,
        DataType::I24 => 2,
        DataType::I32 => 3,
        DataType::F32 => 4,
        DataType::F64 => 5,
    }
}

/// Determine whether a conversion is lossy.
fn is_lossy(settings: &ExportWavSettings, stored_dtype: DataType) -> bool {
    let target = match settings.dtype.as_ref() {
        Some(ExportWavDtype::F32) => DataType::F32,
        Some(ExportWavDtype::I16) => DataType::I16,
        Some(ExportWavDtype::I32) => DataType::I32,
        None => return false, // keep stored type → not lossy
    };
    dtype_rank(&target) < dtype_rank(&stored_dtype)
}

#[cfg(test)]
mod lossy_tests {
    use super::*;

    #[test]
    fn test_is_lossy() {
        let s = |dtype: Option<ExportWavDtype>| ExportWavSettings {
            dtype,
            normalize: false,
            overwrite: false,
            allow_lossy: false,
        };
        // float → int is lossy
        assert!(is_lossy(&s(Some(ExportWavDtype::I16)), DataType::F32));
        assert!(is_lossy(&s(Some(ExportWavDtype::I32)), DataType::F32));
        // int → float is not lossy (higher rank)
        assert!(!is_lossy(&s(Some(ExportWavDtype::F32)), DataType::I16));
        // keep stored → not lossy
        assert!(!is_lossy(&s(None), DataType::F32));
        assert!(!is_lossy(&s(None), DataType::I16));
        // i32 → i16 is lossy
        assert!(is_lossy(&s(Some(ExportWavDtype::I16)), DataType::I32));
    }
}

impl Measurement {
    /// Export measurement data to a WAV file.
    ///
    /// Raw samples are streamed from the HDF5 dataset; sensitivities are
    /// **not** applied.  Use [`ExportWavSettings::normalize`] to scale the
    /// signal instead.
    ///
    /// # Arguments
    ///
    /// * `output_path` - Path for the output `.wav` file (extension added
    ///   automatically if missing).
    /// * `settings` - Export configuration (data type, normalisation, etc.).
    pub fn to_wav(&self, output_path: &Path, settings: ExportWavSettings) -> Result<()> {
        // Validate output path
        let mut out_path = output_path.to_path_buf();
        out_path.set_extension("wav");
        if !settings.overwrite {
            ensure!(
                !out_path.exists(),
                ExportWavFileExistsSnafu { path: out_path }
            );
        }

        let stored_dtype = self.dataType();
        let nchannels = self.nchannels();
        let sr = self.samplerate();

        let out_dtype = match settings.dtype {
            Some(ExportWavDtype::F32) => DataType::F32,
            Some(ExportWavDtype::I16) => DataType::I16,
            Some(ExportWavDtype::I32) => DataType::I32,
            None => stored_dtype,
        };

        // Lossy-check: only when an explicit conversion is requested
        if settings.dtype.is_some() {
            let lossy = is_lossy(&settings, stored_dtype);
            ensure!(
                !lossy || settings.allow_lossy,
                ExportWavLossyConversionSnafu {
                    from: stored_dtype,
                    to: out_dtype,
                }
            );
        }

        let spec = wav_spec(nchannels, sr, &out_dtype);
        let pass_through = out_dtype == stored_dtype && !settings.normalize;

        // Normalization: first pass to find global max absolute value
        let norm_scale: Option<f64> = if settings.normalize {
            let mut max_abs: f64 = 0.0;
            for block in self.raw_iter(None, None, None)? {
                let floats = block.toFloat();
                for &v in floats.iter() {
                    let a = v.abs();
                    if a > max_abs {
                        max_abs = a;
                    }
                }
            }
            Some(if max_abs > 0.0 { 1.0 / max_abs } else { 1.0 })
        } else {
            None
        };

        // Write
        let mut writer = hound::WavWriter::create(&out_path, spec).map_err(|e| {
            ExportWavWriteSnafu {
                msg: format!("Cannot create WAV file: {e}"),
            }
            .build()
        })?;

        for block in self.raw_iter(None, None, None)? {
            if pass_through {
                write_raw_chunk(&mut writer, &block, stored_dtype)?;
            } else {
                let mut floats = block.toFloat();
                if let Some(scale) = norm_scale {
                    floats.mapv_inplace(|v| v * scale);
                }
                write_float_chunk(&mut writer, floats.view(), out_dtype)?;
            }
        }

        writer.finalize().map_err(|e| {
            ExportWavWriteSnafu {
                msg: format!("Error finalizing WAV file: {e}"),
            }
            .build()
        })?;

        Ok(())
    }
}

/// Write a raw chunk directly in its native format (pass-through).
fn write_raw_chunk(
    writer: &mut hound::WavWriter<std::io::BufWriter<std::fs::File>>,
    block: &RawChunk,
    dtype: DataType,
) -> Result<()> {
    use dasp_sample::Sample;
    let wav_err = |e: hound::Error| {
        ExportWavWriteSnafu {
            msg: format!("Error writing WAV sample: {e}"),
        }
        .build()
    };
    macro_rules! write_typed {
        ($arr:expr, $t:ty) => {{
            for &s in $arr.iter() {
                writer.write_sample(s).map_err(&wav_err)?;
            }
        }};
    }
    match (block, dtype) {
        (RawChunk::Datai8(arr), DataType::I8) => write_typed!(arr, i8),
        (RawChunk::Datai16(arr), DataType::I16) => write_typed!(arr, i16),
        (RawChunk::Datai32(arr), DataType::I32) => write_typed!(arr, i32),
        (RawChunk::Dataf32(arr), DataType::F32) => write_typed!(arr, f32),
        (RawChunk::Datai24(arr), DataType::I24) => {
            for &s in arr.iter() {
                writer
                    .write_sample(s.to_sample::<i32>())
                    .map_err(&wav_err)?;
            }
        }
        _ => unreachable!("raw chunk type must match dtype"),
    }
    Ok(())
}

/// Write a float array, converting each sample to the target WAV format
/// via [`dasp_sample::Sample::from_sample`].
fn write_float_chunk(
    writer: &mut hound::WavWriter<std::io::BufWriter<std::fs::File>>,
    floats: ArrayView2<Flt>,
    dtype: DataType,
) -> Result<()> {
    use dasp_sample::{I24, Sample};
    let wav_err = |e: hound::Error| {
        ExportWavWriteSnafu {
            msg: format!("Error writing WAV sample: {e}"),
        }
        .build()
    };
    for &v in floats.iter() {
        let clamped = v.clamp(-1.0, 1.0);
        match dtype {
            DataType::I16 => writer.write_sample(i16::from_sample(clamped)),
            DataType::I24 => writer.write_sample(I24::from_sample(clamped).to_sample::<i32>()),
            DataType::I32 => writer.write_sample(i32::from_sample(clamped)),
            DataType::F32 => writer.write_sample(clamped as f32),
            DataType::F64 => writer.write_sample(clamped as f32),
            _ => unreachable!("unsupported WAV dtype: {dtype:?}"),
        }
        .map_err(&wav_err)?;
    }
    Ok(())
}

#[cfg(test)]
mod tests {
    use super::*;
    use approx::assert_abs_diff_eq;

    /// Create a stereo WAV file with a known sine tone using `hound`,
    /// import it via `Measurement::from_wav`, and verify metadata and
    /// sample data.
    #[test]
    fn test_from_wav() -> anyhow::Result<()> {
        use hound::{SampleFormat, WavSpec, WavWriter};

        let dir = tempfile::tempdir()?;
        let dirpath = dir.path();

        let wav_path = dirpath.join("test_stereo.wav");
        let spec = WavSpec {
            channels: 2,
            sample_rate: 44100,
            bits_per_sample: 32,
            sample_format: SampleFormat::Float,
        };

        {
            let mut writer = WavWriter::create(&wav_path, spec)?;
            for i in 0..100 {
                writer.write_sample(i as f32)?;
                writer.write_sample(100.0 + i as f32)?;
            }
            writer.finalize()?;
        }

        let ch_names = ["pressure", "reference"];
        let _sens = [2.0, 0.5];
        let _qtys = [Qty::AcousticPressure, Qty::Voltage];
        let comment = "Stereo float test";

        let h5_path = dirpath.join("imported_wav");
        let meas = Measurement::from_wav(
            &wav_path,
            Some(h5_path.as_path()),
            Some(&ch_names),
            Some(&_sens),
            Some(&_qtys),
            Some(comment),
            None,
        )?;

        let m = meas.read();
        assert_eq!(m.samplerate(), 44100.0_f64.try_into().unwrap());
        assert_eq!(m.nchannels(), 2);
        assert_eq!(m.channel_names(), ch_names);
        assert_eq!(m.comment(), comment);
        drop(m);

        let h5_filepath = h5_path.with_extension("h5");
        assert!(h5_filepath.exists());

        Ok(())
    }

    /// Round-trip test: import WAV → export WAV → re-import → compare.
    #[test]
    fn test_to_wav_roundtrip() -> anyhow::Result<()> {
        let dir = tempfile::tempdir()?;
        let dirpath = dir.path();

        let wav_in = dirpath.join("original.wav");
        {
            let spec = hound::WavSpec {
                channels: 2,
                sample_rate: 8000,
                bits_per_sample: 32,
                sample_format: hound::SampleFormat::Float,
            };
            let mut writer = hound::WavWriter::create(&wav_in, spec)?;
            for i in 0..64 {
                writer.write_sample(i as f32)?;
                writer.write_sample(-(i as f32))?;
            }
            writer.finalize()?;
        }

        let h5_path = dirpath.join("roundtrip_h5");
        let meas = Measurement::from_wav(&wav_in, Some(&h5_path), None, None, None, None, None)?;
        let wav_out = dirpath.join("exported.wav");
        let settings = ExportWavSettings {
            dtype: None,
            normalize: false,
            overwrite: false,
            allow_lossy: false,
        };
        let original_data = {
            let m = meas.read();
            let data = m.data(None, None, None)?;
            m.to_wav(&wav_out, settings)?;
            data
        };

        let h5_path2 = dirpath.join("reimported_h5");
        let meas2 = Measurement::from_wav(&wav_out, Some(&h5_path2), None, None, None, None, None)?;
        let reimported_data = meas2.read().data(None, None, None)?;

        let orig = original_data.as_slice().expect("original data as slice");
        let reimp = reimported_data
            .as_slice()
            .expect("reimported data as slice");
        assert_eq!(orig.len(), reimp.len());
        for (o, r) in orig.iter().zip(reimp.iter()) {
            assert_abs_diff_eq!(o, r, epsilon = 1e-6);
        }

        Ok(())
    }
}