use alloc::vec;
use alloc::vec::Vec;
use core::f32::consts::PI;
use num_complex::Complex32;
#[cfg(not(feature = "std"))]
use num_traits::Float;
use crate::engine::dsp::downsample::with_default_planner;
use crate::engine::fft::AlignedComplexBuf;
use super::baseband::CENTER_HZ;
use super::demod::{N_SYMBOLS, NSPS_BASEBAND, TONE_SPACING_HZ};
const NFILT: usize = 360;
pub fn subtract_signal_baseband(
idat: &mut [f32],
qdat: &mut [f32],
f0_audio_hz: f32,
shift_baseband: i32,
drift_hz: f32,
channel_symbols: &[u8; N_SYMBOLS],
) {
debug_assert_eq!(idat.len(), qdat.len());
let np = idat.len() as i32;
let nsig = N_SYMBOLS * NSPS_BASEBAND; let f0_baseband_hz = f0_audio_hz - CENTER_HZ;
let mut refi = vec![0.0f32; nsig];
let mut refq = vec![0.0f32; nsig];
let dt = 1.0 / super::baseband::BASEBAND_RATE;
let twopidt = 2.0 * PI * dt;
let mut c = 1.0f32;
let mut s = 0.0f32;
for i in 0..N_SYMBOLS {
let norm = (c * c + s * s).sqrt();
c /= norm;
s /= norm;
let cs = channel_symbols[i] as f32;
let dphi = twopidt
* (f0_baseband_hz
+ (drift_hz / 2.0) * (i as f32 - N_SYMBOLS as f32 / 2.0)
/ (N_SYMBOLS as f32 / 2.0)
+ (cs - 1.5) * TONE_SPACING_HZ);
let (sdphi, cdphi) = dphi.sin_cos();
for j in 0..NSPS_BASEBAND {
let ii = NSPS_BASEBAND * i + j;
refi[ii] = c;
refq[ii] = s;
let (c_next, s_next) = (c * cdphi - s * sdphi, c * sdphi + s * cdphi);
c = c_next;
s = s_next;
}
}
let mut window = [0.0f32; NFILT];
let mut norm = 0.0f32;
for i in 0..NFILT {
window[i] = (PI * i as f32 / (NFILT - 1) as f32).sin();
norm += window[i];
}
for w in window.iter_mut() {
*w /= norm;
}
let mut partial = [0.0f32; NFILT];
for i in 1..NFILT {
partial[i] = partial[i - 1] + window[i];
}
let pad = NFILT;
let nc2 = nsig + 2 * NFILT;
let mut ci = vec![0.0f32; nc2];
let mut cq = vec![0.0f32; nc2];
let i_lo = (1_i64 - shift_baseband as i64).clamp(0, nsig as i64) as usize;
let i_hi = (np as i64 - shift_baseband as i64).clamp(0, nsig as i64) as usize;
for i in i_lo..i_hi {
let k = (shift_baseband + i as i32) as usize;
let id = idat[k];
let qd = qdat[k];
ci[i + pad] = id * refi[i] + qd * refq[i];
cq[i + pad] = qd * refi[i] - id * refq[i];
}
let half = NFILT / 2;
let (cfi, cfq) = lpf_apply_fft(&ci, &cq, &window);
for i in i_lo..i_hi {
let n = if i < half {
partial[half + i]
} else if i > nsig - 1 - half {
partial[half + nsig - 1 - i]
} else {
1.0
};
if n > 0.0 {
let k = (shift_baseband + i as i32) as usize;
let j = i + pad;
idat[k] -= (cfi[j] * refi[i] - cfq[j] * refq[i]) / n;
qdat[k] -= (cfi[j] * refq[i] + cfq[j] * refi[i]) / n;
}
}
}
#[cfg_attr(not(test), allow(dead_code))]
fn lpf_apply_direct(ci: &[f32], cq: &[f32], window: &[f32; NFILT]) -> (Vec<f32>, Vec<f32>) {
let nc2 = ci.len();
debug_assert_eq!(cq.len(), nc2);
let half = NFILT / 2;
let mut cfi = vec![0.0f32; nc2];
let mut cfq = vec![0.0f32; nc2];
for i in half..(nc2 - half) {
let ci_win = &ci[i - half..i - half + NFILT];
let cq_win = &cq[i - half..i - half + NFILT];
let (acc_i, acc_q) = window
.iter()
.zip(ci_win)
.zip(cq_win)
.fold((0.0f32, 0.0f32), |(ai, aq), ((&w, &c_i), &c_q)| {
(ai + w * c_i, aq + w * c_q)
});
cfi[i] = acc_i;
cfq[i] = acc_q;
}
(cfi, cfq)
}
fn lpf_apply_fft(ci: &[f32], cq: &[f32], window: &[f32; NFILT]) -> (Vec<f32>, Vec<f32>) {
const LPF_NFFT: usize = 8192;
let nc2 = ci.len();
debug_assert_eq!(cq.len(), nc2);
let half = NFILT / 2;
let mut cfi = vec![0.0f32; nc2];
let mut cfq = vec![0.0f32; nc2];
if nc2 <= 2 * half {
return (cfi, cfq);
}
let mut kernel_buf = AlignedComplexBuf::zeroed(LPF_NFFT);
let kernel = kernel_buf.as_mut_slice();
for (j, &w) in window.iter().enumerate() {
let d = half as isize - j as isize;
let idx = d.rem_euclid(LPF_NFFT as isize) as usize;
kernel[idx] = Complex32::new(w, 0.0);
}
let fft_fwd = with_default_planner(|planner| planner.plan_forward(LPF_NFFT));
let fft_inv = with_default_planner(|planner| planner.plan_inverse(LPF_NFFT));
fft_fwd.process(kernel);
#[cfg(feature = "fft-rustfft")]
let fac = 1.0f32 / LPF_NFFT as f32;
#[cfg(not(feature = "fft-rustfft"))]
let fac = 1.0f32;
let valid_per_block = LPF_NFFT - NFILT;
let out_lo = half;
let out_hi = nc2 - half;
let mut block_buf = AlignedComplexBuf::zeroed(LPF_NFFT);
let block = block_buf.as_mut_slice();
let mut out_start = out_lo;
while out_start < out_hi {
let block_len = valid_per_block.min(out_hi - out_start);
let in_start = out_start as isize - half as isize;
for (n, b) in block.iter_mut().enumerate() {
let src = in_start + n as isize;
*b = if src >= 0 && (src as usize) < nc2 {
Complex32::new(ci[src as usize], cq[src as usize])
} else {
Complex32::new(0.0, 0.0)
};
}
fft_fwd.process(block);
for (b, k) in block.iter_mut().zip(kernel.iter()) {
*b *= *k;
}
fft_inv.process(block);
for n in 0..block_len {
let v = block[half + n] * fac;
cfi[out_start + n] = v.re;
cfq[out_start + n] = v.im;
}
out_start += block_len;
}
(cfi, cfq)
}
pub fn subtract_all<F>(
idat: &mut [f32],
qdat: &mut [f32],
decodes: &[super::WsprResult],
audio_to_baseband_lag: F,
) where
F: Fn(&super::WsprResult) -> i32,
{
for d in decodes {
let symbols = super::encode_channel_symbols(&d.info_bits);
let f0_audio = d.freq_hz + 1.5 * TONE_SPACING_HZ; let shift_baseband = audio_to_baseband_lag(d);
subtract_signal_baseband(
idat,
qdat,
f0_audio,
shift_baseband,
0.0, &symbols,
);
}
let _ = Vec::<u8>::new(); }
#[cfg(test)]
mod tests {
use super::*;
use crate::wspr::baseband::{NPOINTS_MAX, decimate_to_baseband};
use crate::wspr::tx::synthesize_type1;
#[test]
fn lpf_fft_matches_direct() {
let nsig = super::N_SYMBOLS * super::NSPS_BASEBAND;
let nc2 = nsig + 2 * NFILT;
let mut ci = vec![0.0f32; nc2];
let mut cq = vec![0.0f32; nc2];
let mut state: u32 = 0x1234_5678;
for i in 0..nc2 {
state ^= state << 13;
state ^= state >> 17;
state ^= state << 5;
let noise = (state as f32 / u32::MAX as f32) - 0.5;
let slow = (i as f32 / 4000.0).sin();
let fast = (i as f32 / 25.0).cos();
ci[i] = 0.6 * slow + 0.05 * fast + 0.02 * noise;
cq[i] = 0.4 * slow.cos() - 0.03 * fast + 0.02 * noise;
}
let mut window = [0.0f32; NFILT];
let mut norm = 0.0f32;
for i in 0..NFILT {
window[i] = (PI * i as f32 / (NFILT - 1) as f32).sin();
norm += window[i];
}
for w in window.iter_mut() {
*w /= norm;
}
let (direct_i, direct_q) = lpf_apply_direct(&ci, &cq, &window);
let (fft_i, fft_q) = lpf_apply_fft(&ci, &cq, &window);
assert_eq!(direct_i.len(), fft_i.len());
let half = NFILT / 2;
let mut max_abs_err = 0.0f32;
let mut max_val = 0.0f32;
for i in half..(nc2 - half) {
max_abs_err = max_abs_err
.max((direct_i[i] - fft_i[i]).abs())
.max((direct_q[i] - fft_q[i]).abs());
max_val = max_val.max(direct_i[i].abs()).max(direct_q[i].abs());
}
assert!(
max_abs_err < max_val * 1e-4,
"FFT LPF diverges from direct convolution: max_abs_err={:.3e} max_val={:.3e} (ratio {:.3e})",
max_abs_err,
max_val,
max_abs_err / max_val
);
for i in 0..half {
assert_eq!(direct_i[i], 0.0);
assert_eq!(fft_i[i], 0.0, "fft LPF should leave the pre-half edge at 0");
}
for i in (nc2 - half)..nc2 {
assert_eq!(direct_i[i], 0.0);
assert_eq!(fft_i[i], 0.0, "fft LPF should leave the post-edge at 0");
}
}
#[test]
fn subtract_attenuates_synth_tone() {
let audio = synthesize_type1("K1ABC", "FN42", 37, 12_000, 1500.0, 0.5).expect("synth");
let mut padded = vec![0.0f32; NPOINTS_MAX];
padded[..audio.len()].copy_from_slice(&audio);
let (mut idat, mut qdat) = decimate_to_baseband(&padded);
let pre_pwr: f32 =
idat.iter().map(|&x| x * x).sum::<f32>() + qdat.iter().map(|&x| x * x).sum::<f32>();
let r = crate::wspr::decode_at(&audio, 12_000, 0, 1500.0).expect("decode synth");
let symbols = crate::wspr::encode_channel_symbols(&r.info_bits);
subtract_signal_baseband(
&mut idat,
&mut qdat,
1500.0 + 1.5 * TONE_SPACING_HZ,
0,
0.0,
&symbols,
);
let post_pwr: f32 =
idat.iter().map(|&x| x * x).sum::<f32>() + qdat.iter().map(|&x| x * x).sum::<f32>();
assert!(
post_pwr < pre_pwr * 0.5,
"subtract should remove most of the signal energy: pre={:.2e} post={:.2e}",
pre_pwr,
post_pwr
);
}
}