use fdaf_aec::FdafAec;
use hound::{WavReader, WavWriter, WavSpec};
use clap::Parser;
use std::path::PathBuf;
#[derive(Parser, Debug)]
#[clap(author, version, about, long_about = None)]
struct Args {
#[clap(long, value_parser)]
farend: PathBuf,
#[clap(long, value_parser)]
mic: PathBuf,
#[clap(long, value_parser)]
output: PathBuf,
#[clap(long, value_parser, default_value_t = 0.02)]
step_size: f32,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let args = Args::parse();
println!("--- Running File-Based AEC ---");
println!("- Far-end file: {}", args.farend.display());
println!("- Mic file: {}", args.mic.display());
println!("- Output file: {}", args.output.display());
println!("- Step size: {}", args.step_size);
const FFT_SIZE: usize = 1024;
const FRAME_SIZE: usize = FFT_SIZE / 2;
let (far_signal, spec) = read_wav(&args.farend)?;
let (mic_signal, _) = read_wav(&args.mic)?;
if spec.channels != 1 || spec.sample_rate != 16000 {
eprintln!("Warning: For best results, input WAV files should be mono, 16kHz.");
eprintln!("Current spec: {} channels, {} Hz", spec.channels, spec.sample_rate);
}
let mut aec = FdafAec::new(FFT_SIZE, args.step_size);
let mut processed_signal = Vec::new();
let num_samples = mic_signal.len().min(far_signal.len());
for i in (0..num_samples).step_by(FRAME_SIZE) {
if i + FRAME_SIZE > num_samples { break; }
let far_frame = &far_signal[i..i + FRAME_SIZE];
let mic_frame = &mic_signal[i..i + FRAME_SIZE];
let output_frame = aec.process(far_frame, mic_frame);
processed_signal.extend_from_slice(&output_frame);
}
let mut writer = WavWriter::create(&args.output, spec)?;
for &sample in processed_signal.iter() {
writer.write_sample((sample.clamp(-1.0, 1.0) * i16::MAX as f32) as i16)?;
}
writer.finalize()?;
println!("\nProcessing complete!");
println!("Output saved to '{}'", args.output.display());
Ok(())
}
fn read_wav(path: &PathBuf) -> Result<(Vec<f32>, WavSpec), Box<dyn std::error::Error>> {
let mut reader = WavReader::open(path)?;
let spec = reader.spec();
let max_val = 2_i32.pow(spec.bits_per_sample as u32 - 1) as f32;
let samples = reader.samples::<i32>()
.map(|s| s.unwrap() as f32 / max_val)
.collect();
Ok((samples, spec))
}