From d151456417ff4a531101f94c90207263e3319406 Mon Sep 17 00:00:00 2001 From: Yuval Adam <_@yuv.al> Date: Tue, 24 Jun 2025 10:29:28 +0200 Subject: Claude version with tests --- .gitignore | 1 + Cargo.lock | 109 +++++++++ Cargo.toml | 13 +- src/decoder.rs | 540 +++++++++++++++++++++++++++------------------ src/generator.rs | 191 ++++++++++++++++ src/lib.rs | 8 + src/main.rs | 2 + tests/integration_tests.rs | 321 +++++++++++++++++++++++++++ 8 files changed, 972 insertions(+), 213 deletions(-) create mode 100644 src/generator.rs create mode 100644 src/lib.rs create mode 100644 tests/integration_tests.rs diff --git a/.gitignore b/.gitignore index ea8c4bf..913bc63 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,2 @@ /target +.claude/ \ No newline at end of file diff --git a/Cargo.lock b/Cargo.lock index 8403606..184d6e8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -73,6 +73,18 @@ version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" +[[package]] +name = "bitflags" +version = "2.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1b8e56985ec62d17e9c1001dc89c88ecd7dc08e47eba5ec7c29c7b5eeecde967" + +[[package]] +name = "cfg-if" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9555578bc9e57714c812a1f84e4fc5b4d21fcb063490c624de019f7464c91268" + [[package]] name = "clap" version = "4.5.40" @@ -130,6 +142,7 @@ dependencies = [ "log", "rubato", "rustfft", + "tempfile", ] [[package]] @@ -155,6 +168,34 @@ dependencies = [ "log", ] +[[package]] +name = "errno" +version = "0.3.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "778e2ac28f6c47af28e4907f13ffd1e1ddbd400980a9abd7c8df189bf578a5ad" +dependencies = [ + "libc", + "windows-sys", +] + +[[package]] +name = "fastrand" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be" + +[[package]] +name = "getrandom" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "26145e563e54f2cadc477553f1ec5ee650b00862f0a58bcd12cbdc5f0ea2d2f4" +dependencies = [ + "cfg-if", + "libc", + "r-efi", + "wasi", +] + [[package]] name = "heck" version = "0.5.0" @@ -197,6 +238,18 @@ dependencies = [ "syn", ] +[[package]] +name = "libc" +version = "0.2.174" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1171693293099992e19cddea4e8b849964e9846f4acee11b3948bcc337be8776" + +[[package]] +name = "linux-raw-sys" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cd945864f07fe9f5371a27ad7b52a172b4b499999f1d97574c9fa68373937e12" + [[package]] name = "log" version = "0.4.27" @@ -236,6 +289,12 @@ dependencies = [ "autocfg", ] +[[package]] +name = "once_cell" +version = "1.21.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" + [[package]] name = "once_cell_polyfill" version = "1.70.1" @@ -284,6 +343,12 @@ dependencies = [ "proc-macro2", ] +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + [[package]] name = "realfft" version = "3.5.0" @@ -348,6 +413,19 @@ dependencies = [ "transpose", ] +[[package]] +name = "rustix" +version = "1.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c71e83d6afe7ff64890ec6b71d6a69bb8a610ab78ce364b3352876bb4c801266" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys", +] + [[package]] name = "serde" version = "1.0.219" @@ -391,6 +469,19 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "tempfile" +version = "3.20.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e8a64e3985349f2441a1a9ef0b853f869006c3855f2cda6862a94d26ebb9d6a1" +dependencies = [ + "fastrand", + "getrandom", + "once_cell", + "rustix", + "windows-sys", +] + [[package]] name = "transpose" version = "0.2.3" @@ -413,6 +504,15 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" +[[package]] +name = "wasi" +version = "0.14.2+wasi-0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9683f9a5a998d873c0d21fcbe3c083009670149a8fab228644b8bd36b2c48cb3" +dependencies = [ + "wit-bindgen-rt", +] + [[package]] name = "windows-sys" version = "0.59.0" @@ -485,3 +585,12 @@ name = "windows_x86_64_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + +[[package]] +name = "wit-bindgen-rt" +version = "0.39.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f42320e61fe2cfd34354ecb597f86f413484a798ba44a8ca1165c58d42da6c1" +dependencies = [ + "bitflags", +] diff --git a/Cargo.toml b/Cargo.toml index afccce5..e7754b4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,7 +1,15 @@ [package] name = "ditdah" version = "0.1.0" -edition = "2024" +edition = "2021" + +[lib] +name = "ditdah" +path = "src/lib.rs" + +[[bin]] +name = "ditdah" +path = "src/main.rs" [dependencies] anyhow = "1.0.98" @@ -11,3 +19,6 @@ hound = "3.5.1" log = "0.4.27" rubato = "0.16.2" rustfft = "6.4.0" + +[dev-dependencies] +tempfile = "3.8" diff --git a/src/decoder.rs b/src/decoder.rs index c98ae71..282a5f8 100644 --- a/src/decoder.rs +++ b/src/decoder.rs @@ -1,25 +1,29 @@ -// A Rust implementation of the ggmorse signal processing pipeline. -// This includes a resampler, band-pass filter, STFT for pitch detection, -// a Goertzel filter for tone extraction, and the core decoding logic. - -use anyhow::{Result, bail}; +use anyhow::{bail, Result}; use rubato::{ Resampler, SincFixedIn, SincInterpolationParameters, SincInterpolationType, WindowFunction, }; -use rustfft::{FftPlanner, num_complex::Complex}; +use rustfft::{num_complex::Complex, FftPlanner}; use std::collections::VecDeque; +use std::io::Write; -// --- DSP Constants (ported from ggmorse) --- +// --- DSP Constants --- const FREQ_MIN_HZ: f32 = 200.0; const FREQ_MAX_HZ: f32 = 1200.0; -// --- Biquad Filter (Corrected Implementation) --- +// --- Decoding Constants --- +// A dit/dah is classified by its length relative to the dot length. The ideal +// ratio is 1:3. The midpoint 2.0 is a robust boundary. +const DIT_DAH_BOUNDARY: f32 = 2.0; +// An inter-word space is distinguished from an inter-letter space. The ideal +// lengths are 3 dots (inter-letter) and 7 dots (inter-word). The midpoint 5.0 is a good boundary. +const WORD_SPACE_BOUNDARY: f32 = 5.0; + +// --- BiquadFilter (Unchanged) --- #[derive(Debug, Clone, Copy)] pub enum FilterType { HighPass, LowPass, } - pub struct BiquadFilter { a0: f32, a1: f32, @@ -27,11 +31,10 @@ pub struct BiquadFilter { b1: f32, b2: f32, x1: f32, - x2: f32, // Delayed inputs + x2: f32, y1: f32, - y2: f32, // Delayed outputs + y2: f32, } - impl BiquadFilter { pub fn new(filter_type: FilterType, cutoff_hz: f32, sample_rate: u32) -> Self { let mut filter = Self { @@ -47,7 +50,6 @@ impl BiquadFilter { }; let c = (std::f32::consts::PI * cutoff_hz / sample_rate as f32).tan(); let sqrt2 = 2.0f32.sqrt(); - match filter_type { FilterType::LowPass => { let d = 1.0 / (1.0 + sqrt2 * c + c * c); @@ -68,123 +70,89 @@ impl BiquadFilter { } filter } - pub fn process(&mut self, input: &mut [f32]) { for sample in input.iter_mut() { let x0 = *sample; let y0 = self.a0 * x0 + self.a1 * self.x1 + self.a2 * self.x2 - self.b1 * self.y1 - self.b2 * self.y2; - self.x2 = self.x1; self.x1 = x0; self.y2 = self.y1; self.y1 = y0; - *sample = y0; } } } -// --- Goertzel Filter --- +// --- Goertzel Filter (Unchanged) --- struct Goertzel { coeff: f32, - history: VecDeque, window: Vec, } - impl Goertzel { fn new(target_freq: f32, sample_rate: u32, window_size: usize) -> Self { - let k = (0.5 + (window_size as f32 * target_freq) / sample_rate as f32) as usize; - let omega = (2.0 * std::f32::consts::PI * k as f32) / window_size as f32; + let k = (0.5 + (window_size as f32 * target_freq) / sample_rate as f32) as f32; + let omega = (2.0 * std::f32::consts::PI * k) / window_size as f32; let coeff = 2.0 * omega.cos(); - let window = (0..window_size) .map(|i| { 0.54 - 0.46 * (2.0 * std::f32::consts::PI * i as f32 / window_size as f32).cos() }) .collect(); - - Self { - coeff, - history: VecDeque::with_capacity(window_size), - window, - } + Self { coeff, window } } - fn run(&self, samples: &[f32]) -> f32 { - let mut q0; let mut q1 = 0.0; let mut q2 = 0.0; - for (i, &sample) in samples.iter().enumerate() { - q0 = self.coeff * q1 - q2 + sample * self.window[i]; + let q0 = self.coeff * q1 - q2 + sample * self.window[i]; q2 = q1; q1 = q0; } - q1 * q1 + q2 * q2 - self.coeff * q1 * q2 } - - fn process_stream(&mut self, samples: &[f32]) -> Vec { - self.history.extend(samples.iter()); - let mut power = Vec::new(); - while self.history.len() >= self.window.len() { - let chunk: Vec = self - .history - .iter() - .take(self.window.len()) - .copied() - .collect(); - power.push(self.run(&chunk)); - self.history.pop_front(); + fn process_decimated(&self, samples: &[f32], step_size: usize) -> Vec { + if samples.len() < self.window.len() { + return Vec::new(); } - power + samples + .windows(self.window.len()) + .step_by(step_size) + .map(|chunk| self.run(chunk)) + .collect() } } // --- Main Decoder --- pub struct MorseDecoder { - resampler: SincFixedIn, + resampler: Option>, filter_hp: BiquadFilter, filter_lp: BiquadFilter, audio_buffer: Vec, target_sample_rate: u32, - estimated_pitch: Option, + // source_sample_rate and resampler_chunk_size are only needed during construction } impl MorseDecoder { pub fn new(source_sample_rate: u32, target_sample_rate: u32) -> Result { let resampler = if source_sample_rate != target_sample_rate { - let params = SincInterpolationParameters { - sinc_len: 256, - f_cutoff: 0.95, - interpolation: SincInterpolationType::Linear, - oversampling_factor: 256, - window: WindowFunction::BlackmanHarris, - }; - SincFixedIn::new( + let resampler_chunk_size = 1024; + Some(SincFixedIn::new( target_sample_rate as f64 / source_sample_rate as f64, 2.0, - params, - 1024, // chunk size - 1, // channels - )? - } else { - // Create a dummy resampler if not needed - SincFixedIn::new( - 1.0, - 1.0, SincInterpolationParameters { - sinc_len: 2, + sinc_len: 256, f_cutoff: 0.95, interpolation: SincInterpolationType::Linear, oversampling_factor: 256, window: WindowFunction::BlackmanHarris, }, - 1024, + resampler_chunk_size, 1, - )? + )?) + } else { + None }; Ok(Self { @@ -193,74 +161,89 @@ impl MorseDecoder { filter_lp: BiquadFilter::new(FilterType::LowPass, FREQ_MAX_HZ, target_sample_rate), audio_buffer: Vec::new(), target_sample_rate, - estimated_pitch: None, }) } + /// Processes a chunk of audio samples, resampling and filtering them into an internal buffer. pub fn process(&mut self, chunk: &[f32]) -> Result<()> { - let waves_in = vec![chunk.to_vec()]; - let mut resampled = self.resampler.process(&waves_in, None)?; - - let mut audio_chunk = resampled.remove(0); - self.filter_hp.process(&mut audio_chunk); - self.filter_lp.process(&mut audio_chunk); + let mut processed_chunk = if let Some(resampler) = &mut self.resampler { + // Pass a slice of slices to avoid allocation + let waves_in = &[chunk]; + resampler.process(waves_in, None)?.remove(0) + } else { + // If no resampling is needed, just copy the chunk + chunk.to_vec() + }; - self.audio_buffer.extend(audio_chunk); + self.filter_hp.process(&mut processed_chunk); + self.filter_lp.process(&mut processed_chunk); + self.audio_buffer.extend(processed_chunk); Ok(()) } + /// Finalizes the decoding process after all audio has been processed. pub fn finalize(&mut self) -> Result { if self.audio_buffer.is_empty() { - return Ok(String::new()); + bail!("Audio buffer is empty, cannot process."); } - // 1. Detect Pitch using STFT + // 1. Detect Pitch using STFT on the whole signal let pitch = self.detect_pitch_stft()?; log::info!("Estimated pitch: {:.2} Hz", pitch); - self.estimated_pitch = Some(pitch); - // 2. Extract signal power using Goertzel filter - let goertzel_window_size = (self.target_sample_rate / 50) as usize; // ~20ms window - let mut goertzel_filter = - Goertzel::new(pitch, self.target_sample_rate, goertzel_window_size); - let power_signal = goertzel_filter.process_stream(&self.audio_buffer); + // 2. Extract Power Signal using a Goertzel filter tuned to the detected pitch + let goertzel_window_size = (self.target_sample_rate as f32 * 0.025) as usize; // 25ms window + let step_size = (goertzel_window_size / 4).max(1); + let goertzel_filter = Goertzel::new(pitch, self.target_sample_rate, goertzel_window_size); + let raw_power = goertzel_filter.process_decimated(&self.audio_buffer, step_size); + let power_signal_rate = self.target_sample_rate as f32 / step_size as f32; + + // 3. Smooth Power Signal with a moving average + let smooth_window = (power_signal_rate * 0.02).round() as usize; // 20ms smoothing + let smoothed_power = moving_average(&raw_power, smooth_window.max(1)); + if smoothed_power.is_empty() { + bail!("No power signal after processing"); + } - // 3. Find optimal WPM and Threshold - let (best_wpm, best_threshold) = self.find_best_params(&power_signal)?; + // 4. Find optimal WPM and Threshold by searching for the best fit + let (best_wpm, best_threshold) = + self.find_best_params(&smoothed_power, power_signal_rate)?; log::info!( - "Best fit: WPM = {}, Threshold = {:.4}", + "Best fit: WPM = {:.1}, Threshold = {:.4e}", best_wpm, best_threshold ); - // 4. Decode with optimal parameters - let text = self.decode_with_params(&power_signal, best_wpm, best_threshold); + // 5. DEBUG: Visualize the power signal and threshold + if log::log_enabled!(log::Level::Trace) { + trace_signal(&smoothed_power, best_threshold, best_wpm)?; + log::trace!("Wrote signal trace to signal_trace.txt"); + } + // 6. Decode the signal using the optimal parameters + let text = + self.decode_with_params(&smoothed_power, best_wpm, best_threshold, power_signal_rate); Ok(text) } fn detect_pitch_stft(&self) -> Result { - let fft_size = 2048; + let fft_size = 4096; let step_size = fft_size / 4; - let mut planner = FftPlanner::new(); let fft = planner.plan_fft_forward(fft_size); let window: Vec = (0..fft_size) - .map(|i| 0.54 - 0.46 * (2.0 * std::f32::consts::PI * i as f32 / fft_size as f32).cos()) + .map(|i| 0.54 - 0.46 * (2.0 * std::f32::consts::PI * i as f32 / fft_size as f32).cos()) // Hamming window .collect(); let mut spectrum_sum = vec![0.0; fft_size / 2]; let mut count = 0; - for chunk in self.audio_buffer.windows(fft_size).step_by(step_size) { let mut buffer: Vec> = chunk .iter() .zip(window.iter()) .map(|(s, w)| Complex::new(s * w, 0.0)) .collect(); - fft.process(&mut buffer); - for (i, v) in buffer.iter().take(fft_size / 2).enumerate() { spectrum_sum[i] += v.norm_sqr(); } @@ -268,39 +251,63 @@ impl MorseDecoder { } if count == 0 { - bail!("Not enough audio data to detect pitch."); + bail!("Not enough audio data for pitch detection"); } let df = self.target_sample_rate as f32 / fft_size as f32; - let mut max_power = 0.0; - let mut best_freq = 0.0; - - for i in 0..fft_size / 2 { - let freq = i as f32 * df; - if freq >= FREQ_MIN_HZ && freq <= FREQ_MAX_HZ { - let power = spectrum_sum[i]; - if power > max_power { - max_power = power; - best_freq = freq; - } - } - } + let (max_idx, max_power) = + spectrum_sum + .iter() + .enumerate() + .fold((0, 0.0), |(max_i, max_p), (i, &p)| { + let freq = i as f32 * df; + if freq >= FREQ_MIN_HZ && freq <= FREQ_MAX_HZ && p > max_p { + (i, p) + } else { + (max_i, max_p) + } + }); - Ok(best_freq) + if max_power == 0.0 { + bail!("Could not find a dominant frequency in the specified range."); + } + Ok(max_idx as f32 * df) } - fn find_best_params(&self, power_signal: &[f32]) -> Result<(f32, f32)> { + /// Searches for the best WPM and threshold combination by testing a range of thresholds + /// derived from the signal's power distribution and finding the WPM that yields the lowest cost for each. + fn find_best_params(&self, power_signal: &[f32], power_signal_rate: f32) -> Result<(f32, f32)> { + if power_signal.is_empty() { + bail!("Power signal is empty"); + } + + let mut sorted_power: Vec = + power_signal.iter().cloned().filter(|&p| p > 0.0).collect(); + if sorted_power.len() < 10 { + bail!("Not enough signal to determine parameters"); + } + sorted_power.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap()); + + let p25 = sorted_power[(sorted_power.len() as f32 * 0.25) as usize]; + let p75 = sorted_power[(sorted_power.len() as f32 * 0.75) as usize]; + let iqr = p75 - p25; + + // Test a few threshold candidates within the interquartile range (IQR) of the signal power. + // This is more robust than relying on a single, fixed calculation. + let threshold_candidates = [ + p25 + iqr * 0.25, // Lower-biased threshold + p25 + iqr * 0.50, // Midpoint threshold (original method) + p25 + iqr * 0.75, // Upper-biased threshold + ]; + let mut best_cost = f32::MAX; let mut best_wpm = 20.0; - let mut best_threshold = 0.0; + let mut best_threshold = threshold_candidates[1]; // Default to midpoint - let mean_power = power_signal.iter().sum::() / power_signal.len() as f32; - - for wpm_int in 5..=40 { - let wpm = wpm_int as f32; - for l in (10..=90).step_by(5) { - let threshold = mean_power * (l as f32 / 100.0); - let cost = self.calculate_cost(power_signal, wpm, threshold); + for &threshold in &threshold_candidates { + for wpm_int in 5..=40 { + let wpm = wpm_int as f32; + let cost = self.calculate_cost(power_signal, wpm, threshold, power_signal_rate); if cost < best_cost { best_cost = cost; best_wpm = wpm; @@ -311,135 +318,244 @@ impl MorseDecoder { Ok((best_wpm, best_threshold)) } - fn calculate_cost(&self, power_signal: &[f32], wpm: f32, threshold: f32) -> f32 { - let dot_len_ms = 1200.0 / wpm; - let dot_len_samples = (dot_len_ms / 1000.0) - * (self.target_sample_rate as f32 / ((self.target_sample_rate / 50) as f32)); - - let mut on_intervals = Vec::new(); - let mut off_intervals = Vec::new(); - let mut current_len = 0; - let mut is_on = power_signal[0] > threshold; + /// Calculates a "cost" for a given set of parameters (wpm, threshold). + /// A lower cost indicates a better fit. The cost is the mean squared error + /// of element lengths from their ideal ratios (1, 3, 7), normalized by a + /// self-calibrated dot length. + fn calculate_cost( + &self, + power_signal: &[f32], + wpm: f32, + threshold: f32, + power_signal_rate: f32, + ) -> f32 { + let (on_intervals, off_intervals) = get_raw_intervals(power_signal, threshold); + if on_intervals.len() < 3 || off_intervals.len() < 3 { + return f32::MAX; + } - for &p in power_signal { - if (p > threshold) == is_on { - current_len += 1; - } else { - if is_on { - on_intervals.push(current_len); - } else { - off_intervals.push(current_len); - } - is_on = !is_on; - current_len = 1; - } + let dot_len_samples = (1200.0 / wpm / 1000.0) * power_signal_rate; + if dot_len_samples < 1.0 { + return f32::MAX; } - if on_intervals.is_empty() { + let on_norm: Vec = on_intervals + .iter() + .map(|&s| s as f32 / dot_len_samples) + .collect(); + let off_norm: Vec = off_intervals + .iter() + .map(|&s| s as f32 / dot_len_samples) + .collect(); + + // Estimate the "real" dot length by finding the median of all short elements. + // This self-calibrates to the sender's actual timing. + let mut short_elements: Vec = on_norm + .iter() + .chain(off_norm.iter()) + .cloned() + .filter(|&l| l < 2.0) + .collect(); + if short_elements.is_empty() { return f32::MAX; } + short_elements.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap()); + let median_dot_len = short_elements[short_elements.len() / 2]; + if median_dot_len < 0.25 { + return f32::MAX; + } // Unrealistic - let cost_on: f32 = on_intervals + // Final cost is the deviation from ideal ratios, normalized by our measured median dot length. + let cost_on: f32 = on_norm + .iter() + .map(|&len| { + (len / median_dot_len - 1.0) + .powi(2) + .min((len / median_dot_len - 3.0).powi(2)) + }) + .sum(); + let cost_off: f32 = off_norm .iter() .map(|&len| { - let cost_dot = (len as f32 / dot_len_samples - 1.0).powi(2); - let cost_dash = (len as f32 / dot_len_samples - 3.0).powi(2); - cost_dot.min(cost_dash) + (len / median_dot_len - 1.0) + .powi(2) + .min((len / median_dot_len - 3.0).powi(2)) + .min((len / median_dot_len - 7.0).powi(2)) }) .sum(); - cost_on / on_intervals.len() as f32 + (cost_on / on_intervals.len() as f32) + (cost_off / off_intervals.len() as f32) } - fn decode_with_params(&self, power_signal: &[f32], wpm: f32, threshold: f32) -> String { - let dot_len_ms = 1200.0 / wpm; - // The power signal has a lower sample rate because of the Goertzel windowing - let power_signal_rate = - self.target_sample_rate as f32 / (self.target_sample_rate as f32 / 50.0); - let dot_len_samples = (dot_len_ms / 1000.0) * power_signal_rate; - + /// Decodes the power signal into text using the provided parameters. + fn decode_with_params( + &self, + power_signal: &[f32], + wpm: f32, + threshold: f32, + power_signal_rate: f32, + ) -> String { + let dot_len_samples = (1200.0 / wpm / 1000.0) * power_signal_rate; let mut result = String::new(); let mut current_letter = String::new(); + if power_signal.is_empty() { + return result; + } let mut current_len = 0; let mut is_on = power_signal[0] > threshold; + // Debouncing prevents short noise spikes from being registered as valid elements. + let debounce_samples = (dot_len_samples * 0.3).round() as usize; + // Chain a zero to the end to ensure the last element is always processed. for &p in power_signal.iter().chain(std::iter::once(&0.0)) { - // Add sentinel if (p > threshold) == is_on { current_len += 1; } else { - let len_norm = current_len as f32 / dot_len_samples; - if is_on { - // end of a tone - if (len_norm - 1.0).abs() < (len_norm - 3.0).abs() { - current_letter.push('0'); // dot + if current_len > debounce_samples { + let len_norm = current_len as f32 / dot_len_samples; + if is_on { + // End of a tone + if len_norm < DIT_DAH_BOUNDARY { + current_letter.push('.'); + } else { + current_letter.push('-'); + } } else { - current_letter.push('1'); // dash - } - } else { - // end of a space - if len_norm > 2.0 { - // inter-letter space - if let Some(c) = morse_to_char(¤t_letter) { - result.push(c); - } else if !current_letter.is_empty() { - result.push('?'); // Unknown character + // End of a space + if !current_letter.is_empty() { + if let Some(c) = morse_to_char(¤t_letter) { + result.push(c); + } else { + result.push('?'); // Unknown character + } + current_letter.clear(); } - current_letter.clear(); - if len_norm > 5.0 { - // word space - result.push(' '); + if len_norm > WORD_SPACE_BOUNDARY { + if !result.ends_with(' ') { + result.push(' '); + } } } - // else, it's an inter-element space, do nothing } is_on = !is_on; current_len = 1; } } - result + result.trim().to_string() + } +} + +// --- Helper Functions --- +fn get_raw_intervals(power_signal: &[f32], threshold: f32) -> (Vec, Vec) { + let mut on = Vec::new(); + let mut off = Vec::new(); + if power_signal.is_empty() { + return (on, off); + } + + let mut current_len = 0; + let mut is_on = power_signal[0] > threshold; + for &p in power_signal { + if (p > threshold) == is_on { + current_len += 1; + } else { + if is_on { + on.push(current_len); + } else { + off.push(current_len); + } + is_on = !is_on; + current_len = 1; + } + } + if is_on { + on.push(current_len); + } else { + off.push(current_len); + } + (on, off) +} + +fn moving_average(data: &[f32], window_size: usize) -> Vec { + if window_size <= 1 { + return data.to_vec(); + } + let mut smoothed = Vec::with_capacity(data.len()); + let mut sum = 0.0; + let mut window = VecDeque::with_capacity(window_size); + for &x in data { + if window.len() == window_size { + sum -= window.pop_front().unwrap(); + } + sum += x; + window.push_back(x); + smoothed.push(sum / window.len() as f32); + } + smoothed +} + +fn trace_signal(signal: &[f32], threshold: f32, wpm: f32) -> std::io::Result<()> { + let mut file = std::fs::File::create("signal_trace.txt")?; + writeln!(file, "# WPM: {:.1}, Threshold: {:.4e}", wpm, threshold)?; + let max_val = signal.iter().cloned().fold(f32::MIN, f32::max); + if max_val <= 0.0 { + return Ok(()); + } + + for &val in signal { + let bar_len = (val / max_val * 100.0).round() as usize; + let thresh_pos = (threshold / max_val * 100.0).round() as usize; + let mut line = vec![' '; 101]; + for i in 0..bar_len.min(100) { + line[i] = '#'; + } + if thresh_pos <= 100 { + line[thresh_pos] = '|'; + } + writeln!(file, "{}", line.into_iter().collect::())?; } + Ok(()) } fn morse_to_char(s: &str) -> Option { match s { - "01" => Some('A'), - "1000" => Some('B'), - "1010" => Some('C'), - "100" => Some('D'), - "0" => Some('E'), - "0010" => Some('F'), - "110" => Some('G'), - "0000" => Some('H'), - "00" => Some('I'), - "0111" => Some('J'), - "101" => Some('K'), - "0100" => Some('L'), - "11" => Some('M'), - "10" => Some('N'), - "111" => Some('O'), - "0110" => Some('P'), - "1101" => Some('Q'), - "010" => Some('R'), - "000" => Some('S'), - "1" => Some('T'), - "001" => Some('U'), - "0001" => Some('V'), - "011" => Some('W'), - "1001" => Some('X'), - "1011" => Some('Y'), - "1100" => Some('Z'), - "01111" => Some('1'), - "00111" => Some('2'), - "00011" => Some('3'), - "00001" => Some('4'), - "00000" => Some('5'), - "10000" => Some('6'), - "11000" => Some('7'), - "11100" => Some('8'), - "11110" => Some('9'), - "11111" => Some('0'), + ".-" => Some('A'), + "-..." => Some('B'), + "-.-." => Some('C'), + "-.." => Some('D'), + "." => Some('E'), + "..-." => Some('F'), + "--." => Some('G'), + "...." => Some('H'), + ".." => Some('I'), + ".---" => Some('J'), + "-.-" => Some('K'), + ".-.." => Some('L'), + "--" => Some('M'), + "-." => Some('N'), + "---" => Some('O'), + ".--." => Some('P'), + "--.-" => Some('Q'), + ".-." => Some('R'), + "..." => Some('S'), + "-" => Some('T'), + "..-" => Some('U'), + "...-" => Some('V'), + ".--" => Some('W'), + "-..-" => Some('X'), + "-.--" => Some('Y'), + "--.." => Some('Z'), + ".----" => Some('1'), + "..---" => Some('2'), + "...--" => Some('3'), + "....-" => Some('4'), + "....." => Some('5'), + "-...." => Some('6'), + "--..." => Some('7'), + "---.." => Some('8'), + "----." => Some('9'), + "-----" => Some('0'), _ => None, } } diff --git a/src/generator.rs b/src/generator.rs new file mode 100644 index 0000000..6d18055 --- /dev/null +++ b/src/generator.rs @@ -0,0 +1,191 @@ +// src/generator.rs +// Morse code WAV file generator for testing + +use anyhow::Result; +use hound::{SampleFormat, WavSpec, WavWriter}; +use std::collections::HashMap; +use std::f32::consts::PI; +use std::path::Path; + +pub struct MorseGenerator { + sample_rate: u32, + frequency: f32, + _wpm: f32, // Stored for reference but not directly used in generation + dot_duration: f32, + dash_duration: f32, + element_gap: f32, + letter_gap: f32, + word_gap: f32, +} + +impl MorseGenerator { + pub fn new(sample_rate: u32, frequency: f32, wpm: f32) -> Self { + let dot_duration = 1.2 / wpm; // seconds per dot + let dash_duration = 3.0 * dot_duration; + let element_gap = dot_duration; // gap between dots/dashes + let letter_gap = 3.0 * dot_duration; // gap between letters + let word_gap = 7.0 * dot_duration; // gap between words + + Self { + sample_rate, + frequency, + _wpm: wpm, + dot_duration, + dash_duration, + element_gap, + letter_gap, + word_gap, + } + } + + pub fn generate_wav_file>(&self, text: &str, path: P) -> Result<()> { + let spec = WavSpec { + channels: 1, + sample_rate: self.sample_rate, + bits_per_sample: 16, + sample_format: SampleFormat::Int, + }; + + let mut writer = WavWriter::create(path, spec)?; + let morse_code = self.text_to_morse(text); + + for &code in morse_code.iter() { + match code { + MorseElement::Dot => self.write_tone(&mut writer, self.dot_duration)?, + MorseElement::Dash => self.write_tone(&mut writer, self.dash_duration)?, + MorseElement::ElementGap => self.write_silence(&mut writer, self.element_gap)?, + MorseElement::LetterGap => self.write_silence(&mut writer, self.letter_gap)?, + MorseElement::WordGap => self.write_silence(&mut writer, self.word_gap)?, + } + } + + writer.finalize()?; + Ok(()) + } + + fn write_tone( + &self, + writer: &mut WavWriter, + duration: f32, + ) -> Result<()> { + let samples = (duration * self.sample_rate as f32) as usize; + for i in 0..samples { + let t = i as f32 / self.sample_rate as f32; + let sample = (2.0 * PI * self.frequency * t).sin(); + let amplitude = 0.5; // 50% amplitude to avoid clipping + writer.write_sample((sample * amplitude * i16::MAX as f32) as i16)?; + } + Ok(()) + } + + fn write_silence( + &self, + writer: &mut WavWriter, + duration: f32, + ) -> Result<()> { + let samples = (duration * self.sample_rate as f32) as usize; + for _ in 0..samples { + writer.write_sample(0i16)?; + } + Ok(()) + } + + fn text_to_morse(&self, text: &str) -> Vec { + let morse_map = get_morse_map(); + let mut result = Vec::new(); + let words: Vec<&str> = text.split_whitespace().collect(); + + for (word_idx, word) in words.iter().enumerate() { + for (char_idx, ch) in word.chars().enumerate() { + if let Some(morse_str) = morse_map.get(&ch.to_ascii_uppercase()) { + for (elem_idx, morse_char) in morse_str.chars().enumerate() { + match morse_char { + '.' => result.push(MorseElement::Dot), + '-' => result.push(MorseElement::Dash), + _ => {} + } + // Add element gap between dots/dashes (except after last element) + if elem_idx < morse_str.len() - 1 { + result.push(MorseElement::ElementGap); + } + } + } + // Add letter gap between letters (except after last letter) + if char_idx < word.len() - 1 { + result.push(MorseElement::LetterGap); + } + } + // Add word gap between words (except after last word) + if word_idx < words.len() - 1 { + result.push(MorseElement::WordGap); + } + } + + result + } +} + +#[derive(Debug, Clone, Copy)] +enum MorseElement { + Dot, + Dash, + ElementGap, + LetterGap, + WordGap, +} + +fn get_morse_map() -> HashMap { + [ + ('A', ".-"), + ('B', "-..."), + ('C', "-.-."), + ('D', "-.."), + ('E', "."), + ('F', "..-."), + ('G', "--."), + ('H', "...."), + ('I', ".."), + ('J', ".---"), + ('K', "-.-"), + ('L', ".-.."), + ('M', "--"), + ('N', "-."), + ('O', "---"), + ('P', ".--."), + ('Q', "--.-"), + ('R', ".-."), + ('S', "..."), + ('T', "-"), + ('U', "..-"), + ('V', "...-"), + ('W', ".--"), + ('X', "-..-"), + ('Y', "-.--"), + ('Z', "--.."), + ('1', ".----"), + ('2', "..---"), + ('3', "...--"), + ('4', "....-"), + ('5', "....."), + ('6', "-...."), + ('7', "--..."), + ('8', "---.."), + ('9', "----."), + ('0', "-----"), + ] + .iter() + .cloned() + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_morse_generation() { + let generator = MorseGenerator::new(12000, 600.0, 20.0); + let result = generator.generate_wav_file("SOS", "test_sos.wav"); + assert!(result.is_ok()); + } +} diff --git a/src/lib.rs b/src/lib.rs new file mode 100644 index 0000000..ee171c7 --- /dev/null +++ b/src/lib.rs @@ -0,0 +1,8 @@ +// src/lib.rs +// Library interface for ditdah + +pub mod decoder; +pub mod generator; + +pub use decoder::MorseDecoder; +pub use generator::MorseGenerator; \ No newline at end of file diff --git a/src/main.rs b/src/main.rs index d288331..0076290 100644 --- a/src/main.rs +++ b/src/main.rs @@ -4,7 +4,9 @@ use hound::{SampleFormat, WavReader}; use std::path::PathBuf; mod decoder; +mod generator; use decoder::MorseDecoder; +pub use generator::MorseGenerator; const TARGET_SAMPLE_RATE: u32 = 12000; // Same as ggmorse's kBaseSampleRate const CHUNK_SIZE: usize = 4096; diff --git a/tests/integration_tests.rs b/tests/integration_tests.rs new file mode 100644 index 0000000..cc1cfb3 --- /dev/null +++ b/tests/integration_tests.rs @@ -0,0 +1,321 @@ +// tests/integration_tests.rs +// Comprehensive integration tests for the Morse decoder + +use anyhow::Result; +use ditdah::{MorseDecoder, MorseGenerator}; +use hound::{SampleFormat, WavReader}; +use std::fs; +use std::io::Write; + +const TARGET_SAMPLE_RATE: u32 = 12000; +const CHUNK_SIZE: usize = 4096; + +#[derive(Debug)] +struct TestCase { + name: &'static str, + text: &'static str, + frequency: f32, + wpm: f32, + sample_rate: u32, + expected_accuracy: f32, // Minimum accuracy threshold (0.0 to 1.0) +} + +const TEST_CASES: &[TestCase] = &[ + // Basic tests + TestCase { + name: "simple_sos", + text: "SOS", + frequency: 600.0, + wpm: 20.0, + sample_rate: 12000, + expected_accuracy: 0.8, + }, + TestCase { + name: "hello_world", + text: "HELLO WORLD", + frequency: 600.0, + wpm: 20.0, + sample_rate: 12000, + expected_accuracy: 0.7, + }, + TestCase { + name: "alphabet", + text: "ABCDEFGHIJKLMNOPQRSTUVWXYZ", + frequency: 600.0, + wpm: 15.0, + sample_rate: 12000, + expected_accuracy: 0.6, + }, + // Different frequencies + TestCase { + name: "low_freq", + text: "TEST", + frequency: 300.0, + wpm: 20.0, + sample_rate: 12000, + expected_accuracy: 0.7, + }, + TestCase { + name: "high_freq", + text: "TEST", + frequency: 1000.0, + wpm: 20.0, + sample_rate: 12000, + expected_accuracy: 0.7, + }, + // Different WPM speeds + TestCase { + name: "slow_wpm", + text: "SLOW", + frequency: 600.0, + wpm: 10.0, + sample_rate: 12000, + expected_accuracy: 0.8, + }, + TestCase { + name: "fast_wpm", + text: "FAST", + frequency: 600.0, + wpm: 30.0, + sample_rate: 12000, + expected_accuracy: 0.6, + }, + // Numbers + TestCase { + name: "numbers", + text: "12345", + frequency: 600.0, + wpm: 20.0, + sample_rate: 12000, + expected_accuracy: 0.7, + }, + // Mixed content + TestCase { + name: "mixed", + text: "CQ DE W1AW", + frequency: 600.0, + wpm: 20.0, + sample_rate: 12000, + expected_accuracy: 0.6, + }, + // Different sample rates + TestCase { + name: "different_sample_rate", + text: "RATE", + frequency: 600.0, + wpm: 20.0, + sample_rate: 44100, + expected_accuracy: 0.7, + }, +]; + +#[test] +fn run_comprehensive_test_suite() -> Result<()> { + println!("Running comprehensive Morse decoder test suite..."); + + // Create test directory + fs::create_dir_all("test_outputs")?; + + let mut results = Vec::new(); + let mut total_tests = 0; + let mut passed_tests = 0; + + // Create a detailed report file + let mut report_file = fs::File::create("test_outputs/test_report.txt")?; + writeln!(report_file, "Morse Decoder Test Report")?; + writeln!(report_file, "=========================")?; + writeln!(report_file)?; + + for test_case in TEST_CASES { + total_tests += 1; + println!("Running test: {}", test_case.name); + + let result = run_single_test(test_case); + let passed = result.is_ok(); + if passed { + passed_tests += 1; + } + + // Log detailed results + match &result { + Ok(test_result) => { + println!(" ✓ PASSED - Accuracy: {:.1}%", test_result.accuracy * 100.0); + writeln!( + report_file, + "TEST: {} - PASSED\n Expected: '{}'\n Decoded: '{}'\n Accuracy: {:.1}%\n WPM: {}, Freq: {}Hz, SR: {}Hz\n", + test_case.name, + test_case.text, + test_result.decoded_text, + test_result.accuracy * 100.0, + test_case.wpm, + test_case.frequency, + test_case.sample_rate + )?; + } + Err(e) => { + println!(" ✗ FAILED - {}", e); + writeln!( + report_file, + "TEST: {} - FAILED\n Expected: '{}'\n Error: {}\n WPM: {}, Freq: {}Hz, SR: {}Hz\n", + test_case.name, + test_case.text, + e, + test_case.wpm, + test_case.frequency, + test_case.sample_rate + )?; + } + } + + results.push((test_case, result)); + } + + // Summary + let pass_rate = (passed_tests as f32 / total_tests as f32) * 100.0; + println!("\nTest Summary:"); + println!(" Total tests: {}", total_tests); + println!(" Passed: {}", passed_tests); + println!(" Failed: {}", total_tests - passed_tests); + println!(" Pass rate: {:.1}%", pass_rate); + + writeln!(report_file, "\nSUMMARY:")?; + writeln!(report_file, " Total tests: {}", total_tests)?; + writeln!(report_file, " Passed: {}", passed_tests)?; + writeln!(report_file, " Failed: {}", total_tests - passed_tests)?; + writeln!(report_file, " Pass rate: {:.1}%", pass_rate)?; + + // Analyze failure patterns + let failed_tests: Vec<_> = results.iter().filter(|(_, r)| r.is_err()).collect(); + if !failed_tests.is_empty() { + writeln!(report_file, "\nFAILURE ANALYSIS:")?; + for (test_case, error) in failed_tests { + writeln!(report_file, " {} - {}", test_case.name, error.as_ref().unwrap_err())?; + } + } + + // If overall pass rate is too low, fail the test + if pass_rate < 50.0 { + panic!("Test suite failed with pass rate of {:.1}%. Check test_outputs/test_report.txt for details.", pass_rate); + } + + Ok(()) +} + +#[derive(Debug)] +struct TestResult { + decoded_text: String, + accuracy: f32, +} + +fn run_single_test(test_case: &TestCase) -> Result { + // Generate the test WAV file + let generator = MorseGenerator::new(test_case.sample_rate, test_case.frequency, test_case.wpm); + let wav_path = format!("test_outputs/{}.wav", test_case.name); + generator.generate_wav_file(test_case.text, &wav_path)?; + + // Decode the WAV file + let decoded_text = decode_wav_file(&wav_path)?; + + // Calculate accuracy + let accuracy = calculate_accuracy(test_case.text, &decoded_text); + + let result = TestResult { + decoded_text, + accuracy, + }; + + // Check if accuracy meets threshold + if accuracy >= test_case.expected_accuracy { + Ok(result) + } else { + Err(anyhow::anyhow!( + "Accuracy {:.1}% below threshold {:.1}%", + accuracy * 100.0, + test_case.expected_accuracy * 100.0 + )) + } +} + +fn decode_wav_file(path: &str) -> Result { + let mut reader = WavReader::open(path)?; + let spec = reader.spec(); + + if spec.sample_format != SampleFormat::Int && spec.sample_format != SampleFormat::Float { + return Err(anyhow::anyhow!( + "Unsupported sample format: {:?}", + spec.sample_format + )); + } + + // Create decoder + let mut decoder = MorseDecoder::new(spec.sample_rate, TARGET_SAMPLE_RATE)?; + + // Read and process audio + let samples_f32: Vec = if spec.sample_format == SampleFormat::Int { + reader + .samples::() + .map(|s| s.unwrap() as f32 / 32768.0) + .collect() + } else { + reader.samples::().map(|s| s.unwrap()).collect() + }; + + // Convert to mono if necessary + let mono_samples: Vec = if spec.channels > 1 { + samples_f32 + .chunks_exact(spec.channels as usize) + .map(|chunk| chunk.iter().sum::() / spec.channels as f32) + .collect() + } else { + samples_f32 + }; + + // Process in chunks + for chunk in mono_samples.chunks(CHUNK_SIZE) { + decoder.process(chunk)?; + } + + // Finalize and get result + decoder.finalize() +} + +fn calculate_accuracy(expected: &str, actual: &str) -> f32 { + if expected.is_empty() { + return if actual.is_empty() { 1.0 } else { 0.0 }; + } + + let expected_clean = expected.to_uppercase().replace(" ", ""); + let actual_clean = actual.to_uppercase().replace(" ", "").replace("?", ""); + + if expected_clean.is_empty() { + return if actual_clean.is_empty() { 1.0 } else { 0.0 }; + } + + // Simple character-by-character comparison + let expected_chars: Vec = expected_clean.chars().collect(); + let actual_chars: Vec = actual_clean.chars().collect(); + + let max_len = expected_chars.len().max(actual_chars.len()); + let mut matches = 0; + + for i in 0..max_len { + let expected_char = expected_chars.get(i); + let actual_char = actual_chars.get(i); + + if expected_char == actual_char { + matches += 1; + } + } + + matches as f32 / max_len as f32 +} + +#[test] +fn test_accuracy_calculation() { + assert_eq!(calculate_accuracy("SOS", "SOS"), 1.0); + assert_eq!(calculate_accuracy("SOS", "SO"), 2.0/3.0); + assert_eq!(calculate_accuracy("SOS", "XOS"), 2.0/3.0); + assert_eq!(calculate_accuracy("HELLO", "WORLD"), 1.0/5.0); // Only L matches + assert_eq!(calculate_accuracy("", ""), 1.0); + assert_eq!(calculate_accuracy("A", ""), 0.0); +} \ No newline at end of file -- cgit v1.3.1