use denoize::audio::{
ensure_memory_limit, estimate_audio_working_set_bytes, estimate_file_memory_bytes,
estimate_stream_memory_bytes, read_audio, read_wav_bytes, write_wav_bytes,
write_wav_channel_mask_to_file, WavStreamReader, WavStreamWriter,
};
use denoize::denoiser::{DenoiserConfig, Preset, ProcessingMode};
use denoize::service::{self, BackendChoice, ProcessingOptions};
use denoize::{
AacEncoder, Algorithm, AtomicOutput, Backend, BackendOptions, ChannelMode, CommitMode,
DownmixMode, EncodeOptions, OnnxModelConfig, SgmseProfile, StreamingDenoiser, WindowType,
};
use serde::{Deserialize, Serialize};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Instant;
const VERSION: &str = env!("CARGO_PKG_VERSION");
const STREAM_BLOCK_FRAMES: usize = 8192;
static CANCELLED: AtomicBool = AtomicBool::new(false);
static CANCEL_HANDLER: OnceLock<Result<(), String>> = OnceLock::new();
#[derive(Serialize)]
struct ProcessResultJson<'a> {
input: &'a str,
output: &'a str,
backend: &'a str,
channels: usize,
frames: usize,
sample_rate: u32,
elapsed_ms: f64,
}
#[derive(Serialize)]
struct StreamResultJson<'a> {
input: &'a str,
output: &'a str,
backend: &'static str,
channels: u16,
frames: usize,
sample_rate: u32,
stream: bool,
}
#[derive(Serialize)]
#[serde(tag = "event", rename_all = "lowercase")]
enum BatchJson<'a> {
Progress {
status: &'a str,
completed: usize,
total: usize,
elapsed_seconds: f64,
eta_seconds: f64,
input: &'a str,
},
Summary {
total: usize,
succeeded: usize,
skipped: usize,
failed: usize,
cancelled: bool,
output: &'a str,
},
}
fn serialize_json_line<T: Serialize + ?Sized>(payload: &T) -> String {
serde_json::to_string(payload).expect("fixed CLI JSON payload must serialize")
}
fn round_to_three_decimals(value: f64) -> f64 {
format!("{value:.3}")
.parse()
.expect("formatted JSON number must parse")
}
fn process_result_json_line(
input: &str,
output: &str,
backend: &str,
channels: usize,
frames: usize,
sample_rate: u32,
elapsed_ms: f64,
) -> String {
serialize_json_line(&ProcessResultJson {
input,
output,
backend,
channels,
frames,
sample_rate,
elapsed_ms: round_to_three_decimals(elapsed_ms),
})
}
fn stream_result_json_line(
input: &str,
output: &str,
channels: u16,
frames: usize,
sample_rate: u32,
) -> String {
serialize_json_line(&StreamResultJson {
input,
output,
backend: "classical",
channels,
frames,
sample_rate,
stream: true,
})
}
fn batch_progress_json_line(
status: &str,
completed: usize,
total: usize,
elapsed_seconds: f64,
eta_seconds: f64,
input: &str,
) -> String {
serialize_json_line(&BatchJson::Progress {
status,
completed,
total,
elapsed_seconds: round_to_three_decimals(elapsed_seconds),
eta_seconds: round_to_three_decimals(eta_seconds),
input,
})
}
fn batch_summary_json_line(
total: usize,
succeeded: usize,
skipped: usize,
failed: usize,
cancelled: bool,
output: &str,
) -> String {
serialize_json_line(&BatchJson::Summary {
total,
succeeded,
skipped,
failed,
cancelled,
output,
})
}
fn install_cancel_handler() -> Result<(), String> {
CANCEL_HANDLER
.get_or_init(|| {
ctrlc::set_handler(|| CANCELLED.store(true, Ordering::SeqCst))
.map_err(|error| format!("install Ctrl+C handler: {error}"))
})
.clone()
}
fn usage() -> String {
let backends = Backend::available_names().join("|");
format!(
"\
denoize {VERSION} — pure-Rust audio denoiser engineered for the world's highest sound quality
Classical DSP + optional AI backends (RNNoise, DeepFilterNet v3, MP-SENet, BSRNN).
Input: WAV/BWF/RF64, AIFF, CAF, FLAC, Ogg Opus/Vorbis, MP3, M4A/ALAC, AAC (built in; no ffmpeg).
Output: WAV, FLAC, Ogg Opus, MP3, M4A, AAC.
USAGE:
denoize <INPUT> <OUTPUT.wav|flac|opus|ogg|mp3|m4a|aac> [OPTIONS]
denoize live [--input-device NAME] [--output-device NAME] [OPTIONS]
denoize live --list-devices
denoize models <COMMAND> [MODEL|all] [OPTIONS] (run `denoize models --help`)
denoize metrics <REFERENCE> <TEST> [--json|--markdown]
denoize compare <CLEAN> <NOISY> <ENHANCED> [--json|--html]
OPTIONS:
--config <PATH> load TOML defaults (CLI options take precedence)
-b, --backend <NAME> auto|{backends} (default: classical)
-a, --algorithm <NAME> omlsa|logmmse|mmse|wiener|specsub|specsub-nl|specsub-geo
-p, --preset <NAME> speech|music|aggressive|gentle|restore|hifi
--mode <NAME> speech|music|ambient processing intent
-s, --strength <0..1> denoising strength (default: 0.6)
--profile <MS> learn noise from first MS ms (default: auto-detect)
--no-profile no profiling; rely on blind IMCRA bootstrap
--no-adapt freeze the noise estimate
--adaptive-noise learn noise from noise-only regions throughout the file
--vad speech-aware segmentation and silence suppression
--frame <N> FFT size: 512|1024|2048|4096|8192 (default: 2048)
--overlap <F> overlap ratio 0.5..0.95 (default: 0.75)
--window <NAME> hann|hamming|sine|blackman|kaiser|flattop|dpss
--kaiser-beta <B> Kaiser window beta (default: 8.0)
--dpss-nw <NW> DPSS time-bandwidth product (default: 3.0)
--multiband enable multiband spectral subtraction
--perceptual enable Bark-scale perceptual gain weighting
--postfilter enable musical-noise suppression post-filter
--smoothing <0..1> gain release smoothing (default: 0.6)
--makeup <DB> makeup gain in dB (default: 0.0)
--no-dc-block disable DC-blocking pre-filter
--quality <LEVEL> high|ultra
--no-transient disable transient/onset protection
--cepstral enable cepstral gain smoothing
--no-cepstral disable cepstral smoothing
--pre-emphasis enable pre/de-emphasis
--no-pre-emphasis disable pre-emphasis
--report print settings report and exit
--mp3-bitrate <KBPS> MP3 CBR bitrate (default: 192)
--m4a-bitrate <KBPS> M4A/AAC CBR bitrate (default: 192)
--aac-encoder <NAME> oxide|fdk (default: oxide)
--downmix <MODE> preserve|stereo (default: preserve; lossy outputs reject surround unless explicit)
--loudness <LUFS> normalize integrated loudness after denoising
--true-peak <DBTP> true-peak ceiling with --loudness (default: -1)
--onnx-model <PATH> waveform ONNX model (required for -b onnx)
--onnx-rate <HZ> ONNX model sample rate (default: 16000)
--channels <MODE> independent|linked|mid-side (default: independent)
--sgmse-profile <P> fast|balanced|quality (default: balanced)
--deterministic serialize processing for reproducible audio output
--seed <N> SGMSE sampler seed (implies --deterministic)
--batch process files in INPUT directory into OUTPUT directory
--stream bounded-memory classical WAV-to-WAV processing
--stream-frames <N> streaming block size in frames (default: 8192)
--max-memory <MB> refuse inputs whose estimated working set exceeds MB
--recursive include subdirectories in batch mode
--jobs <N> concurrent batch workers (default: CPU count)
--output-format <EXT> convert every batch output to this format
--force allow replacing existing output files
--resume skip completed files recorded by batch state
--no-progress suppress batch progress and ETA output
--json emit a machine-readable result
--no-metadata do not copy input tags/artwork/chapters to the output
--input-device <NAME> live capture device (default: system default)
--output-device <NAME> live playback device (default: system default)
--chunk-ms <MS> live processing chunk duration (default: 100)
-h, --help show this help
-V, --version show version
BACKENDS (build with --features full for all):
classical Enhanced STFT/IMCRA/OMLSA pipeline (default)
rnnoise RNNoise via nnnoiseless (requires --features rnnoise)
deepfilter DeepFilterNet v3 (requires --features deepfilter)
onnx External waveform ONNX model (requires --features onnx)
mpsenet MP-SENet magnitude/phase model (requires --features mpsenet)
bsrnn ESPnet BSRNN spectral model (requires --features bsrnn)
mossformer2 ClearerVoice MossFormer2 model (requires --features mossformer2)
sgmse SGMSE+ diffusion model (requires --features sgmse)
gtcrn Official low-complexity streaming GTCRN (requires --features gtcrn)
PRESETS:
hifi Flagship transparency: OMLSA + protections + advanced DSP
speech Voice-optimised balance
music Instruments; enables perceptual + postfilter
"
)
}
#[derive(Clone, Debug, Default)]
struct Overrides {
backend: Option<Backend>,
auto_backend: bool,
algorithm: Option<Algorithm>,
preset: Option<Preset>,
mode: Option<ProcessingMode>,
strength: Option<f64>,
profile_ms: Option<f64>,
no_profile: bool,
no_adapt: bool,
adaptive_noise: bool,
vad: bool,
frame_size: Option<usize>,
overlap: Option<f64>,
window: Option<WindowType>,
kaiser_beta: Option<f64>,
dpss_nw: Option<f64>,
multiband: bool,
perceptual: bool,
postfilter: bool,
smoothing: Option<f64>,
makeup: Option<f64>,
no_dc_block: bool,
report: bool,
quality: Option<String>,
no_transient: bool,
cepstral: bool,
no_cepstral: bool,
pre_emphasis: bool,
no_pre_emphasis: bool,
mp3_bitrate_kbps: Option<u32>,
m4a_bitrate_kbps: Option<u32>,
aac_encoder: Option<AacEncoder>,
downmix: Option<DownmixMode>,
loudness_lufs: Option<f64>,
true_peak_dbtp: Option<f64>,
onnx_model: Option<String>,
onnx_sample_rate: Option<u32>,
channel_mode: Option<ChannelMode>,
sgmse_profile: Option<SgmseProfile>,
deterministic: bool,
seed: Option<u64>,
batch: bool,
stream: bool,
stream_frames: Option<usize>,
max_memory_mb: Option<usize>,
recursive: bool,
jobs: Option<usize>,
output_format: Option<String>,
force: bool,
resume: bool,
no_progress: bool,
json: bool,
no_metadata: bool,
input_device: Option<String>,
output_device: Option<String>,
chunk_ms: Option<u32>,
list_devices: bool,
}
#[derive(Debug, Default, Deserialize)]
#[serde(default, deny_unknown_fields)]
struct FileConfig {
backend: Option<String>,
algorithm: Option<String>,
preset: Option<String>,
mode: Option<String>,
strength: Option<f64>,
profile_ms: Option<f64>,
adaptive_noise: bool,
vad: bool,
frame_size: Option<usize>,
overlap: Option<f64>,
window: Option<String>,
smoothing: Option<f64>,
makeup_db: Option<f64>,
quality: Option<String>,
mp3_bitrate_kbps: Option<u32>,
m4a_bitrate_kbps: Option<u32>,
aac_encoder: Option<String>,
loudness_lufs: Option<f64>,
true_peak_dbtp: Option<f64>,
onnx_model: Option<String>,
onnx_rate: Option<u32>,
channels: Option<String>,
sgmse_profile: Option<String>,
downmix: Option<String>,
deterministic: bool,
seed: Option<u64>,
batch: bool,
stream: bool,
stream_frames: Option<usize>,
max_memory_mb: Option<usize>,
recursive: bool,
jobs: Option<usize>,
output_format: Option<String>,
force: bool,
resume: bool,
progress: Option<bool>,
preserve_metadata: Option<bool>,
}
fn load_config(path: &str) -> Result<Overrides, String> {
let source = std::fs::read_to_string(path)
.map_err(|error| format!("failed to read config {path}: {error}"))?;
parse_config(&source, path)
}
fn parse_config(source: &str, path: &str) -> Result<Overrides, String> {
let config: FileConfig =
toml::from_str(source).map_err(|error| format!("invalid config {path}: {error}"))?;
let mut ov = Overrides::default();
if let Some(name) = config.backend {
if name.eq_ignore_ascii_case("auto") {
ov.auto_backend = true;
} else {
ov.backend = Some(
Backend::parse(&name)
.ok_or_else(|| format!("unknown backend in config: {name}"))?,
);
}
}
if let Some(name) = config.algorithm {
ov.algorithm = Some(
Algorithm::parse(&name)
.ok_or_else(|| format!("unknown algorithm in config: {name}"))?,
);
}
if let Some(name) = config.preset {
ov.preset =
Some(Preset::parse(&name).ok_or_else(|| format!("unknown preset in config: {name}"))?);
}
if let Some(name) = config.mode {
ov.mode = Some(
ProcessingMode::parse(&name)
.ok_or_else(|| format!("unknown mode in config: {name}"))?,
);
}
if let Some(name) = config.window {
ov.window = Some(
WindowType::parse(&name).ok_or_else(|| format!("unknown window in config: {name}"))?,
);
}
if let Some(name) = config.channels {
ov.channel_mode = Some(
ChannelMode::parse(&name)
.ok_or_else(|| format!("unknown channel mode in config: {name}"))?,
);
}
if let Some(name) = config.downmix {
ov.downmix = Some(DownmixMode::parse(&name).ok_or_else(|| {
format!("unknown downmix mode in config: {name} (expected preserve or stereo)")
})?);
}
if let Some(name) = config.aac_encoder {
ov.aac_encoder = Some(AacEncoder::parse(&name).ok_or_else(|| {
format!("unknown AAC encoder in config: {name} (expected oxide or fdk)")
})?);
}
if let Some(profile) = config.sgmse_profile {
ov.sgmse_profile = Some(SgmseProfile::parse(&profile).ok_or_else(|| {
format!(
"unknown SGMSE profile in config: {profile} (expected fast, balanced, or quality)"
)
})?);
}
ov.strength = config.strength;
ov.profile_ms = config.profile_ms;
ov.adaptive_noise = config.adaptive_noise;
ov.vad = config.vad;
ov.frame_size = config.frame_size;
ov.overlap = config.overlap;
ov.smoothing = config.smoothing;
ov.makeup = config.makeup_db;
ov.quality = config.quality.map(|value| value.to_ascii_lowercase());
ov.mp3_bitrate_kbps = config.mp3_bitrate_kbps;
ov.m4a_bitrate_kbps = config.m4a_bitrate_kbps;
ov.loudness_lufs = config.loudness_lufs;
ov.true_peak_dbtp = if config.loudness_lufs.is_none() && config.true_peak_dbtp == Some(-1.0) {
None
} else {
config.true_peak_dbtp
};
ov.onnx_model = config.onnx_model;
ov.onnx_sample_rate = config.onnx_rate;
ov.deterministic = config.deterministic;
ov.seed = config.seed;
if ov.seed.is_some() {
ov.deterministic = true;
}
ov.batch = config.batch;
ov.stream = config.stream;
ov.stream_frames = config.stream_frames;
ov.max_memory_mb = config.max_memory_mb;
ov.recursive = config.recursive;
ov.jobs = config.jobs;
ov.output_format = config.output_format;
ov.force = config.force;
ov.resume = config.resume;
ov.no_progress = config.progress == Some(false);
ov.no_metadata = config.preserve_metadata == Some(false);
validate_resource_options(&ov)?;
Ok(ov)
}
fn parse_value<T>(args: &[String], i: &mut usize, flag: &str) -> Result<T, String>
where
T: std::str::FromStr,
<T as std::str::FromStr>::Err: std::fmt::Display,
{
*i += 1;
if *i >= args.len() {
return Err(format!("missing value for {flag}"));
}
args[*i]
.parse::<T>()
.map_err(|e| format!("invalid value for {flag}: {e}"))
}
fn parse_args(args: &[String]) -> Result<(String, String, Overrides), String> {
let mut input: Option<String> = None;
let mut output: Option<String> = None;
let config_path = args
.windows(2)
.find(|pair| pair[0] == "--config")
.map(|pair| pair[1].as_str());
if args.last().map(String::as_str) == Some("--config") {
return Err("missing value for --config".into());
}
let mut ov = match config_path {
Some(path) => load_config(path)?,
None => Overrides::default(),
};
let mut i = 0;
while i < args.len() {
let a = &args[i];
match a.as_str() {
"--config" => {
let _: String = parse_value(args, &mut i, a)?;
}
"-h" | "--help" => {
print!("{}", usage());
std::process::exit(0);
}
"-V" | "--version" => {
println!("denoize {VERSION}");
std::process::exit(0);
}
"-b" | "--backend" => {
let name: String = parse_value(args, &mut i, a)?;
if name.eq_ignore_ascii_case("auto") {
ov.auto_backend = true;
ov.backend = None;
i += 1;
continue;
}
ov.auto_backend = false;
ov.backend = Some(Backend::parse(&name).ok_or_else(|| {
format!(
"unknown backend: {name} (available: {:?})",
Backend::available_names()
)
})?);
}
"-a" | "--algorithm" => {
let name: String = parse_value(args, &mut i, a)?;
ov.algorithm = Some(
Algorithm::parse(&name).ok_or_else(|| format!("unknown algorithm: {name}"))?,
);
}
"-p" | "--preset" => {
let name: String = parse_value(args, &mut i, a)?;
ov.preset =
Some(Preset::parse(&name).ok_or_else(|| format!("unknown preset: {name}"))?);
}
"--mode" => {
let name: String = parse_value(args, &mut i, a)?;
ov.mode = Some(ProcessingMode::parse(&name).ok_or_else(|| {
format!("unknown mode: {name} (expected speech, music, or ambient)")
})?);
}
"-s" | "--strength" => ov.strength = Some(parse_value(args, &mut i, a)?),
"--profile" => ov.profile_ms = Some(parse_value(args, &mut i, a)?),
"--no-profile" => ov.no_profile = true,
"--no-adapt" => ov.no_adapt = true,
"--adaptive-noise" => ov.adaptive_noise = true,
"--vad" => ov.vad = true,
"--frame" => ov.frame_size = Some(parse_value(args, &mut i, a)?),
"--overlap" => ov.overlap = Some(parse_value(args, &mut i, a)?),
"--window" => {
let name: String = parse_value(args, &mut i, a)?;
ov.window = Some(
WindowType::parse(&name).ok_or_else(|| format!("unknown window: {name}"))?,
);
}
"--kaiser-beta" => ov.kaiser_beta = Some(parse_value(args, &mut i, a)?),
"--dpss-nw" => ov.dpss_nw = Some(parse_value(args, &mut i, a)?),
"--multiband" => ov.multiband = true,
"--perceptual" => ov.perceptual = true,
"--postfilter" => ov.postfilter = true,
"--smoothing" => ov.smoothing = Some(parse_value(args, &mut i, a)?),
"--makeup" => ov.makeup = Some(parse_value(args, &mut i, a)?),
"--no-dc-block" => ov.no_dc_block = true,
"--report" => ov.report = true,
"--quality" => {
let q: String = parse_value(args, &mut i, a)?;
ov.quality = Some(q.to_ascii_lowercase());
}
"--no-transient" => ov.no_transient = true,
"--cepstral" => ov.cepstral = true,
"--no-cepstral" => ov.no_cepstral = true,
"--pre-emphasis" => ov.pre_emphasis = true,
"--no-pre-emphasis" => ov.no_pre_emphasis = true,
"--mp3-bitrate" => ov.mp3_bitrate_kbps = Some(parse_value(args, &mut i, a)?),
"--m4a-bitrate" => ov.m4a_bitrate_kbps = Some(parse_value(args, &mut i, a)?),
"--aac-encoder" => {
let name: String = parse_value(args, &mut i, a)?;
ov.aac_encoder = Some(AacEncoder::parse(&name).ok_or_else(|| {
format!("unknown AAC encoder: {name} (expected oxide or fdk)")
})?);
}
"--downmix" => {
let mode: String = parse_value(args, &mut i, a)?;
ov.downmix = Some(DownmixMode::parse(&mode).ok_or_else(|| {
format!("unknown downmix mode: {mode} (expected preserve or stereo)")
})?);
}
"--loudness" => ov.loudness_lufs = Some(parse_value(args, &mut i, a)?),
"--true-peak" => ov.true_peak_dbtp = Some(parse_value(args, &mut i, a)?),
"--onnx-model" => ov.onnx_model = Some(parse_value(args, &mut i, a)?),
"--onnx-rate" => ov.onnx_sample_rate = Some(parse_value(args, &mut i, a)?),
"--channels" => {
let mode: String = parse_value(args, &mut i, a)?;
ov.channel_mode = Some(ChannelMode::parse(&mode).ok_or_else(|| {
format!(
"unknown channel mode: {mode} (expected independent, linked, or mid-side)"
)
})?);
}
"--sgmse-profile" => {
let profile: String = parse_value(args, &mut i, a)?;
ov.sgmse_profile = Some(SgmseProfile::parse(&profile).ok_or_else(|| {
format!(
"unknown SGMSE profile: {profile} (expected fast, balanced, or quality)"
)
})?);
}
"--deterministic" => ov.deterministic = true,
"--seed" => {
ov.seed = Some(parse_value(args, &mut i, a)?);
ov.deterministic = true;
}
"--batch" => ov.batch = true,
"--stream" => ov.stream = true,
"--stream-frames" => ov.stream_frames = Some(parse_value(args, &mut i, a)?),
"--max-memory" => ov.max_memory_mb = Some(parse_value(args, &mut i, a)?),
"--recursive" => ov.recursive = true,
"--jobs" => ov.jobs = Some(parse_value(args, &mut i, a)?),
"--output-format" => ov.output_format = Some(parse_value(args, &mut i, a)?),
"--force" => ov.force = true,
"--resume" => ov.resume = true,
"--no-progress" => ov.no_progress = true,
"--json" => ov.json = true,
"--no-metadata" => ov.no_metadata = true,
"--input-device" => ov.input_device = Some(parse_value(args, &mut i, a)?),
"--output-device" => ov.output_device = Some(parse_value(args, &mut i, a)?),
"--chunk-ms" => ov.chunk_ms = Some(parse_value(args, &mut i, a)?),
"--list-devices" => ov.list_devices = true,
"-" => {
if input.is_none() {
input = Some(a.clone());
} else if output.is_none() {
output = Some(a.clone());
} else {
return Err("unexpected extra argument: -".into());
}
}
other if other.starts_with('-') => {
return Err(format!("unknown option: {other}"));
}
_ => {
if input.is_none() {
input = Some(a.clone());
} else if output.is_none() {
output = Some(a.clone());
} else {
return Err(format!("unexpected extra argument: {a}"));
}
}
}
i += 1;
}
let input = input.ok_or("missing INPUT")?;
let output = output.ok_or("missing OUTPUT audio path")?;
Ok((input, output, ov))
}
fn validate_resource_options(ov: &Overrides) -> Result<(), String> {
if ov.max_memory_mb == Some(0) {
return Err("--max-memory must be at least 1 MiB".into());
}
if ov.stream_frames == Some(0) {
return Err("--stream-frames must be at least 1 frame".into());
}
Ok(())
}
fn build_config(ov: &Overrides, sample_rate: u32) -> DenoiserConfig {
let mut cfg = match ov.preset {
Some(p) => p.config(sample_rate),
None => DenoiserConfig::default(sample_rate),
};
if let Some(mode) = ov.mode {
mode.apply(&mut cfg);
}
if let Some(a) = ov.algorithm {
cfg.algorithm = a;
}
if let Some(s) = ov.strength {
cfg.strength = s;
}
if ov.no_profile {
cfg.profile_ms = -1.0;
} else if let Some(ms) = ov.profile_ms {
cfg.profile_ms = ms;
}
if ov.no_adapt {
cfg.adapt = false;
}
if ov.adaptive_noise {
cfg.adaptive_noise = true;
}
if ov.vad {
cfg.vad = true;
}
if let Some(f) = ov.frame_size {
cfg.frame_size = f;
}
if let Some(o) = ov.overlap {
cfg.overlap = o;
}
if let Some(w) = ov.window {
cfg.window = w;
}
if let Some(b) = ov.kaiser_beta {
cfg.window_params.kaiser_beta = b;
}
if let Some(nw) = ov.dpss_nw {
cfg.window_params.dpss_bandwidth = nw;
}
if ov.multiband {
cfg.multiband = true;
}
if ov.perceptual {
cfg.perceptual_weighting = true;
}
if ov.postfilter {
cfg.musical_noise_postfilter = true;
}
if let Some(s) = ov.smoothing {
cfg.smoothing = s;
}
if let Some(m) = ov.makeup {
cfg.makeup_gain_db = m;
}
if ov.no_dc_block {
cfg.dc_block = false;
}
if let Some(ref q) = ov.quality {
match q.as_str() {
"high" => {
if cfg.frame_size < 2048 {
cfg.frame_size = 2048;
}
if cfg.overlap < 0.8 {
cfg.overlap = 0.8;
}
cfg.transient_protect = true;
cfg.cepstral_smoothing = true;
cfg.perceptual_weighting = true;
cfg.musical_noise_postfilter = true;
if !ov.no_pre_emphasis {
cfg.pre_emphasis = true;
}
}
"ultra" | "max" | "highest" => {
cfg.frame_size = cfg.frame_size.max(4096);
cfg.overlap = 0.875;
cfg.window = WindowType::Kaiser;
cfg.window_params.kaiser_beta = 10.0;
cfg.transient_protect = true;
cfg.cepstral_smoothing = true;
cfg.perceptual_weighting = true;
cfg.musical_noise_postfilter = true;
cfg.pre_emphasis = true;
if ov.strength.is_none() && cfg.strength > 0.4 {
cfg.strength = 0.32;
}
}
_ => {}
}
}
if ov.no_transient {
cfg.transient_protect = false;
}
if ov.cepstral {
cfg.cepstral_smoothing = true;
}
if ov.no_cepstral {
cfg.cepstral_smoothing = false;
}
if ov.pre_emphasis {
cfg.pre_emphasis = true;
}
if ov.no_pre_emphasis {
cfg.pre_emphasis = false;
}
cfg
}
fn print_report(input: &str, audio: &denoize::Audio, cfg: &DenoiserConfig, backend: Backend) {
let hop = (cfg.frame_size as f64 * (1.0 - cfg.overlap)).round() as usize;
let g_min_db = -20.0 - 25.0 * cfg.strength;
let dur = audio.frames() as f64 / audio.sample_rate as f64;
println!("input : {input}");
println!(
"format : {}ch, {:.2}s ({} frames), {} Hz, {}-bit {:?}",
audio.channels(),
dur,
audio.frames(),
audio.sample_rate,
audio.bits_per_sample,
audio.sample_format,
);
println!("layout : {}", audio.channel_layout());
if let Some(mask) = audio.channel_mask {
println!("mask : {mask}");
}
if let Some(pan) = audio.pan_info() {
let positions = pan
.iter()
.enumerate()
.map(|(index, info)| {
format!(
"ch{}={:.0}°/{:.0}°",
index + 1,
info.azimuth_degrees,
info.elevation_degrees
)
})
.collect::<Vec<_>>()
.join(", ");
println!("pan : {positions}");
}
println!("backend : {backend:?}");
println!("algorithm : {:?}", cfg.algorithm);
println!(
"strength : {:.2} (gain floor ~{:.0} dB)",
cfg.strength, g_min_db
);
println!(
"STFT : frame={}, hop={}, overlap={:.0}%, window={:?}",
cfg.frame_size,
hop,
cfg.overlap * 100.0,
cfg.window,
);
println!(
"advanced : multiband={}, perceptual={}, postfilter={}",
cfg.multiband, cfg.perceptual_weighting, cfg.musical_noise_postfilter
);
println!("smoothing : {:.2}", cfg.smoothing);
println!(
"profile : {}",
if cfg.profile_ms < 0.0 {
"disabled".to_string()
} else if cfg.profile_ms == 0.0 {
"auto (leading silence)".to_string()
} else {
format!("{:.0} ms", cfg.profile_ms)
}
);
println!("adapt : {}", cfg.adapt);
println!("adaptive-profile: {}", cfg.adaptive_noise);
println!("dc-block : {}", cfg.dc_block);
println!("makeup : {:.1} dB", cfg.makeup_gain_db);
println!(
"hi-fi : transient={}, cepstral={}, pre-emphasis={}",
cfg.transient_protect, cfg.cepstral_smoothing, cfg.pre_emphasis
);
}
fn run(args: &[String]) -> Result<(), String> {
if args.first().map(String::as_str) == Some("live") {
return run_live(&args[1..]);
}
if args.first().map(String::as_str) == Some("models") {
return run_models(&args[1..]);
}
if args.first().map(String::as_str) == Some("metrics") {
return run_metrics(&args[1..]);
}
if args.first().map(String::as_str) == Some("compare") {
return run_compare(&args[1..]);
}
let (input, output, ov) = parse_args(args)?;
validate_resource_options(&ov)?;
if ov.batch {
if ov.stream {
return Err("--stream cannot be combined with --batch".into());
}
return run_batch(&input, &output, &ov);
}
if ov.stream {
return run_streaming_wav(&input, &output, ov);
}
run_one(&input, &output, ov)
}
#[cfg(feature = "live")]
fn run_live(args: &[String]) -> Result<(), String> {
let mut parseable = vec!["-".to_string(), "-".to_string()];
parseable.extend_from_slice(args);
let (_, _, ov) = parse_args(&parseable)?;
validate_resource_options(&ov)?;
if ov.list_devices {
let (inputs, outputs) = denoize::live::device_names()?;
println!("Input devices:");
for device in inputs {
println!(" {device}");
}
println!("Output devices:");
for device in outputs {
println!(" {device}");
}
return Ok(());
}
let backend = if ov.auto_backend {
service::select_live_backend()
} else {
ov.backend.unwrap_or(Backend::Classical)
};
let sample_rate = 48_000;
let denoiser = build_config(&ov, sample_rate);
let backend_options = BackendOptions {
onnx: ov.onnx_model.map(|path| OnnxModelConfig {
path: path.into(),
sample_rate: ov.onnx_sample_rate.unwrap_or(16_000),
}),
channel_mode: ov.channel_mode.unwrap_or_default(),
sgmse_profile: ov.sgmse_profile.unwrap_or_default(),
deterministic: ov.deterministic,
seed: ov.seed,
};
denoize::live::run(denoize::live::LiveConfig {
input_device: ov.input_device,
output_device: ov.output_device,
chunk_ms: ov.chunk_ms.unwrap_or(100),
backend,
backend_options,
denoiser,
})
}
#[cfg(not(feature = "live"))]
fn run_live(_args: &[String]) -> Result<(), String> {
Err("live audio is unavailable in this build; rebuild with --features live".into())
}
fn ensure_output_available(path: &std::path::Path, force: bool) -> Result<(), String> {
if force {
return Ok(());
}
match std::fs::symlink_metadata(path) {
Ok(_) => Err(format!(
"output already exists: {} (use --force to replace it)",
path.display()
)),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(format!(
"inspect output destination {}: {error}",
path.display()
)),
}
}
fn run_one(input: &str, output: &str, ov: Overrides) -> Result<(), String> {
validate_resource_options(&ov)?;
if input != "-" {
let estimate = estimate_file_memory_bytes(std::path::Path::new(input))?;
ensure_memory_limit(estimate, ov.max_memory_mb, "input preflight")?;
}
let metadata = if input != "-" && !ov.no_metadata {
denoize::metadata::read_extended(std::path::Path::new(input))?
} else {
None
};
let mut audio = if input == "-" {
let mut bytes = Vec::new();
std::io::Read::read_to_end(&mut std::io::stdin(), &mut bytes)
.map_err(|error| format!("failed to read stdin: {error}"))?;
let estimate = (bytes.len() as u64).saturating_mul(8).max(1024 * 1024);
ensure_memory_limit(estimate, ov.max_memory_mb, "stdin input preflight")?;
read_wav_bytes(bytes)?
} else {
read_audio(input)?
};
ensure_memory_limit(
estimate_audio_working_set_bytes(&audio),
ov.max_memory_mb,
"decoded audio working set",
)?;
let cfg = build_config(&ov, audio.sample_rate);
let backend_choice = if ov.auto_backend {
BackendChoice::Auto
} else {
BackendChoice::Explicit(ov.backend.unwrap_or(Backend::Classical))
};
let backend = service::select_backend(
backend_choice,
audio.frames() as f64 / audio.sample_rate as f64,
ov.quality.as_deref(),
);
if ov.auto_backend && !ov.json {
eprintln!(
"denoize: auto-selected backend {}",
service::backend_name(backend)
);
}
if ov.report {
print_report(input, &audio, &cfg, backend);
return Ok(());
}
if output != "-" {
ensure_output_available(std::path::Path::new(output), ov.force)?;
}
let mut enc = EncodeOptions::default();
if let Some(kbps) = ov.mp3_bitrate_kbps {
enc.mp3_bitrate_kbps = kbps;
}
if let Some(kbps) = ov.m4a_bitrate_kbps {
enc.m4a_bitrate_bps = kbps.saturating_mul(1000);
}
if let Some(encoder) = ov.aac_encoder {
enc.aac_encoder = encoder;
}
if let Some(downmix) = ov.downmix {
enc.downmix = downmix;
}
let backend_options = BackendOptions {
onnx: ov.onnx_model.map(|path| OnnxModelConfig {
path: path.into(),
sample_rate: ov.onnx_sample_rate.unwrap_or(16_000),
}),
channel_mode: ov.channel_mode.unwrap_or_default(),
sgmse_profile: ov.sgmse_profile.unwrap_or_default(),
deterministic: ov.deterministic,
seed: ov.seed,
};
let result = service::process_audio(
&mut audio,
ProcessingOptions {
backend: backend_choice,
quality: ov.quality.clone(),
denoiser: cfg,
backend_options,
loudness_lufs: ov.loudness_lufs,
true_peak_dbtp: ov.true_peak_dbtp.unwrap_or(-1.0),
},
)?;
if let Some(report) = result.loudness {
if !ov.json {
eprintln!(
"denoize: loudness {:.2} -> {:.2} LUFS, true peak {:.2} dBTP, gain {:+.2} dB",
report.input_lufs, report.output_lufs, report.true_peak_dbtp, report.gain_db
);
}
} else if ov.true_peak_dbtp.is_some() {
return Err("--true-peak requires --loudness".into());
}
if output == "-" {
let bytes = write_wav_bytes(&audio)?;
std::io::Write::write_all(&mut std::io::stdout(), &bytes)
.map_err(|error| format!("failed to write stdout: {error}"))?;
} else {
let output_path = std::path::Path::new(output);
denoize::write_audio_transactional(
output_path,
&audio,
enc,
metadata,
if ov.force {
CommitMode::Replace
} else {
CommitMode::NoClobber
},
)?;
if ov.json {
println!(
"{}",
process_result_json_line(
input,
output,
service::backend_name(result.backend),
audio.channels(),
audio.frames(),
audio.sample_rate,
result.elapsed.as_secs_f64() * 1_000.0,
)
);
}
}
Ok(())
}
fn run_streaming_wav(input: &str, output: &str, ov: Overrides) -> Result<(), String> {
validate_resource_options(&ov)?;
if input == "-" || output == "-" {
return Err("--stream requires filesystem WAV input and output paths".into());
}
let input_path = std::path::Path::new(input);
let output_path = std::path::Path::new(output);
let is_wav = |path: &std::path::Path| {
path.extension()
.and_then(|extension| extension.to_str())
.map(|extension| extension.eq_ignore_ascii_case("wav"))
.unwrap_or(false)
};
if !is_wav(input_path) || !is_wav(output_path) {
return Err("--stream currently supports WAV-to-WAV paths only".into());
}
ensure_output_available(output_path, ov.force)?;
if ov.auto_backend
|| ov
.backend
.is_some_and(|backend| backend != Backend::Classical)
{
return Err("--stream currently supports only the classical backend".into());
}
if ov
.channel_mode
.is_some_and(|mode| mode != ChannelMode::Independent)
{
return Err("--stream requires independent channels".into());
}
if ov.vad || ov.loudness_lufs.is_some() || ov.true_peak_dbtp.is_some() {
return Err("--stream does not support VAD or loudness normalization".into());
}
let metadata = if !ov.no_metadata {
denoize::metadata::read_extended(input_path)?
} else {
None
};
let mut reader = WavStreamReader::open(input_path)?;
let spec = reader.spec();
let channel_mask = reader.channel_mask();
let cfg = build_config(&ov, spec.sample_rate);
let block_frames = ov.stream_frames.unwrap_or(STREAM_BLOCK_FRAMES);
ensure_memory_limit(
estimate_stream_memory_bytes(
spec.channels as usize,
block_frames,
cfg.frame_size,
spec.sample_rate,
),
ov.max_memory_mb,
"streaming working set",
)?;
if cfg.vad {
return Err("--stream does not support VAD; omit --mode speech or --vad".into());
}
if ov.report {
println!(
"input : {input}\nformat : {}ch, {} Hz, {}-bit {:?}\nbackend : classical\nstream : enabled ({} frames/block)",
spec.channels, spec.sample_rate, spec.bits_per_sample, spec.sample_format, block_frames
);
return Ok(());
}
let mut transaction = AtomicOutput::new(output_path)?;
let frames = (|| -> Result<usize, String> {
let mut processor = StreamingDenoiser::new(cfg, spec.channels as usize)?;
let sink = std::io::BufWriter::new(transaction.file_mut());
let mut writer = WavStreamWriter::from_sink(sink, spec)?;
let mut frames = 0usize;
while let Some(block) = reader.next_block(block_frames)? {
let block_frames = block.first().map(Vec::len).unwrap_or(0);
let enhanced = processor.process_block(&block)?;
writer.write_block(&enhanced)?;
frames = frames.saturating_add(block_frames);
}
let tail = processor.finish()?;
writer.write_block(&tail)?;
writer.finalize()?;
Ok(frames)
})()?;
write_wav_channel_mask_to_file(transaction.file_mut(), spec.channels as usize, channel_mask)?;
if let Some(metadata) = metadata {
denoize::metadata::write_extended_to_file(metadata, transaction.file_mut())?;
}
transaction.commit(if ov.force {
CommitMode::Replace
} else {
CommitMode::NoClobber
})?;
if ov.json {
println!(
"{}",
stream_result_json_line(input, output, spec.channels, frames, spec.sample_rate)
);
} else {
eprintln!(
"denoize: streaming classical WAV complete: {}ch x {} frames",
spec.channels, frames
);
}
Ok(())
}
enum BatchFileOutcome {
Completed,
Skipped,
}
#[derive(Debug, Default, PartialEq, Eq)]
struct BatchCounts {
succeeded: usize,
skipped: usize,
failed: usize,
}
fn count_batch_results<E>(results: &[Result<BatchFileOutcome, E>]) -> BatchCounts {
let mut counts = BatchCounts::default();
for result in results {
match result {
Ok(BatchFileOutcome::Completed) => counts.succeeded += 1,
Ok(BatchFileOutcome::Skipped) => counts.skipped += 1,
Err(_) => counts.failed += 1,
}
}
counts
}
fn run_batch(input: &str, output: &str, ov: &Overrides) -> Result<(), String> {
use rayon::prelude::*;
let input_dir = std::path::Path::new(input);
let output_dir = std::path::Path::new(output);
if !input_dir.is_dir() {
return Err(format!("batch input is not a directory: {input}"));
}
if let Some(jobs) = ov.jobs {
if jobs == 0 {
return Err("--jobs must be at least 1".into());
}
}
let output_extension = ov
.output_format
.as_deref()
.map(normalize_output_extension)
.transpose()?;
std::fs::create_dir_all(output_dir).map_err(|e| format!("create batch output: {e}"))?;
install_cancel_handler()?;
CANCELLED.store(false, Ordering::SeqCst);
let files = collect_batch_files(input_dir, ov.recursive)?;
if files.is_empty() {
return Err("batch input contains no supported audio files".into());
}
if let Some(extension) = output_extension {
let mut destinations = std::collections::HashSet::new();
for path in &files {
let relative = path.strip_prefix(input_dir).map_err(|e| e.to_string())?;
let mut destination = output_dir.join(relative);
destination.set_extension(extension);
if !destinations.insert(destination.clone()) {
return Err(format!(
"multiple inputs map to the same batch output: {}",
destination.display()
));
}
}
}
let state_path = output_dir.join(".denoize-state");
let completed_paths = if ov.resume {
read_batch_state(&state_path)?
} else {
std::collections::HashSet::new()
};
let state_file = if ov.resume {
Some(Arc::new(Mutex::new(
std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&state_path)
.map_err(|error| format!("open resume state {}: {error}", state_path.display()))?,
)))
} else {
None
};
let finished = AtomicUsize::new(0);
let started = Instant::now();
let process_file = |path: &std::path::PathBuf| -> Result<BatchFileOutcome, String> {
if CANCELLED.load(Ordering::SeqCst) {
return Err("cancelled".into());
}
let relative = path.strip_prefix(input_dir).map_err(|e| e.to_string())?;
let mut destination = output_dir.join(relative);
if let Some(extension) = output_extension {
destination.set_extension(extension);
}
let state_key = relative.to_string_lossy().replace('\\', "/");
if ov.resume && completed_paths.contains(&state_key) && destination.is_file() {
report_batch_progress(&finished, files.len(), started, path, "skipped", ov);
return Ok(BatchFileOutcome::Skipped);
}
if let Some(parent) = destination.parent() {
std::fs::create_dir_all(parent)
.map_err(|e| format!("create {}: {e}", parent.display()))?;
}
let mut options = ov.clone();
options.batch = false;
options.json = false;
run_one(
&path.to_string_lossy(),
&destination.to_string_lossy(),
options,
)?;
if let Some(file) = &state_file {
use std::io::Write;
let mut file = file.lock().map_err(|_| "resume state lock poisoned")?;
writeln!(file, "{state_key}")
.map_err(|error| format!("write resume state: {error}"))?;
file.flush()
.map_err(|error| format!("flush resume state: {error}"))?;
}
report_batch_progress(&finished, files.len(), started, path, "completed", ov);
Ok(BatchFileOutcome::Completed)
};
let results = if ov.deterministic {
files.iter().map(process_file).collect::<Vec<_>>()
} else if let Some(jobs) = ov.jobs {
rayon::ThreadPoolBuilder::new()
.num_threads(jobs)
.build()
.map_err(|e| format!("create batch worker pool: {e}"))?
.install(|| files.par_iter().map(process_file).collect::<Vec<_>>())
} else {
files.par_iter().map(process_file).collect::<Vec<_>>()
};
let counts = count_batch_results(&results);
let failures: Vec<_> = results
.iter()
.filter_map(|result| result.as_ref().err())
.collect();
debug_assert_eq!(counts.failed, failures.len());
if ov.json {
println!(
"{}",
batch_summary_json_line(
files.len(),
counts.succeeded,
counts.skipped,
counts.failed,
CANCELLED.load(Ordering::SeqCst),
output,
)
);
} else {
eprintln!(
"denoize: batch complete: {} succeeded ({} skipped), {} failed",
counts.succeeded, counts.skipped, counts.failed
);
for error in &failures {
eprintln!("denoize: batch error: {error}");
}
}
if failures.is_empty() {
Ok(())
} else {
Err(format!("{} batch file(s) failed", failures.len()))
}
}
fn read_batch_state(path: &std::path::Path) -> Result<std::collections::HashSet<String>, String> {
match std::fs::read_to_string(path) {
Ok(source) => Ok(source
.lines()
.filter(|line| !line.is_empty())
.map(str::to_owned)
.collect()),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(Default::default()),
Err(error) => Err(format!("read resume state {}: {error}", path.display())),
}
}
fn report_batch_progress(
finished: &AtomicUsize,
total: usize,
started: Instant,
path: &std::path::Path,
status: &str,
ov: &Overrides,
) {
let count = finished.fetch_add(1, Ordering::Relaxed) + 1;
let elapsed = started.elapsed().as_secs_f64();
let eta = if count == 0 {
0.0
} else {
elapsed / count as f64 * total.saturating_sub(count) as f64
};
if ov.json {
let input = path.to_string_lossy();
println!(
"{}",
batch_progress_json_line(status, count, total, elapsed, eta, input.as_ref())
);
} else if !ov.no_progress {
eprintln!(
"denoize: batch {count}/{total} {status} {} ({elapsed:.1}s elapsed, ETA {eta:.1}s)",
path.display()
);
}
}
fn collect_batch_files(
root: &std::path::Path,
recursive: bool,
) -> Result<Vec<std::path::PathBuf>, String> {
let mut pending = vec![root.to_path_buf()];
let mut files = Vec::new();
while let Some(directory) = pending.pop() {
for entry in std::fs::read_dir(&directory)
.map_err(|e| format!("read batch input {}: {e}", directory.display()))?
{
let path = entry.map_err(|e| format!("read batch entry: {e}"))?.path();
if path.is_dir() && recursive {
pending.push(path);
} else if path.is_file() && is_supported_audio_path(&path) {
files.push(path);
}
}
}
files.sort();
Ok(files)
}
fn is_supported_audio_path(path: &std::path::Path) -> bool {
path.extension()
.and_then(|extension| extension.to_str())
.map(|extension| {
matches!(
extension.to_ascii_lowercase().as_str(),
"wav"
| "rf64"
| "bwf"
| "aif"
| "aiff"
| "aifc"
| "caf"
| "mp3"
| "m4a"
| "mp4"
| "aac"
| "flac"
| "opus"
| "ogg"
| "oga"
| "vorbis"
)
})
.unwrap_or(false)
}
fn normalize_output_extension(value: &str) -> Result<&str, String> {
let extension = value.trim_start_matches('.');
if matches!(
extension.to_ascii_lowercase().as_str(),
"wav" | "mp3" | "m4a" | "aac" | "flac" | "opus" | "ogg"
) {
Ok(extension)
} else {
Err(format!("unsupported --output-format: {value}"))
}
}
#[cfg(test)]
mod json_output_tests {
use super::*;
use serde_json::Value;
const SPECIAL_INPUT: &str = "input-cafe\u{301}-quote\"-slash\\-line\n-control\u{1}.wav";
const SPECIAL_OUTPUT: &str = "output-cafe\u{301}-quote\"-slash\\-line\n-control\u{2}.wav";
fn parse_json_line(line: &str) -> Value {
assert!(
!line.contains("\\u{"),
"Rust escape leaked into JSON: {line}"
);
assert!(
!line.contains('\n'),
"serialized JSON line contains a physical newline"
);
serde_json::from_str(line).expect("CLI output must be valid JSON")
}
#[test]
fn process_result_json_round_trips_special_paths() {
let value = parse_json_line(&process_result_json_line(
SPECIAL_INPUT,
SPECIAL_OUTPUT,
"classical",
2,
48_001,
48_000,
1.2345,
));
assert_eq!(value.as_object().unwrap().len(), 7);
assert_eq!(value["input"].as_str(), Some(SPECIAL_INPUT));
assert_eq!(value["output"].as_str(), Some(SPECIAL_OUTPUT));
assert_eq!(value["backend"].as_str(), Some("classical"));
assert_eq!(value["channels"].as_u64(), Some(2));
assert_eq!(value["frames"].as_u64(), Some(48_001));
assert_eq!(value["sample_rate"].as_u64(), Some(48_000));
assert_eq!(value["elapsed_ms"].as_f64(), Some(1.234));
}
#[test]
fn stream_result_json_round_trips_special_paths() {
let value = parse_json_line(&stream_result_json_line(
SPECIAL_INPUT,
SPECIAL_OUTPUT,
2,
8_193,
44_100,
));
assert_eq!(value.as_object().unwrap().len(), 7);
assert_eq!(value["input"].as_str(), Some(SPECIAL_INPUT));
assert_eq!(value["output"].as_str(), Some(SPECIAL_OUTPUT));
assert_eq!(value["backend"].as_str(), Some("classical"));
assert_eq!(value["channels"].as_u64(), Some(2));
assert_eq!(value["frames"].as_u64(), Some(8_193));
assert_eq!(value["sample_rate"].as_u64(), Some(44_100));
assert_eq!(value["stream"].as_bool(), Some(true));
}
#[test]
fn batch_progress_json_round_trips_special_paths() {
let value = parse_json_line(&batch_progress_json_line(
"completed",
3,
5,
1.23456,
0.45678,
SPECIAL_INPUT,
));
assert_eq!(value.as_object().unwrap().len(), 7);
assert_eq!(value["event"].as_str(), Some("progress"));
assert_eq!(value["status"].as_str(), Some("completed"));
assert_eq!(value["completed"].as_u64(), Some(3));
assert_eq!(value["total"].as_u64(), Some(5));
assert_eq!(value["elapsed_seconds"].as_f64(), Some(1.235));
assert_eq!(value["eta_seconds"].as_f64(), Some(0.457));
assert_eq!(value["input"].as_str(), Some(SPECIAL_INPUT));
}
#[test]
fn batch_summary_json_round_trips_special_paths() {
let value = parse_json_line(&batch_summary_json_line(7, 4, 2, 1, false, SPECIAL_OUTPUT));
assert_eq!(value.as_object().unwrap().len(), 7);
assert_eq!(value["event"].as_str(), Some("summary"));
assert_eq!(value["total"].as_u64(), Some(7));
assert_eq!(value["succeeded"].as_u64(), Some(4));
assert_eq!(value["skipped"].as_u64(), Some(2));
assert_eq!(value["failed"].as_u64(), Some(1));
assert_eq!(value["cancelled"].as_bool(), Some(false));
assert_eq!(value["output"].as_str(), Some(SPECIAL_OUTPUT));
}
}
#[cfg(test)]
mod batch_tests {
use super::*;
fn temporary_directory() -> std::path::PathBuf {
std::env::temp_dir().join(format!(
"denoize-batch-test-{}-{:?}",
std::process::id(),
std::thread::current().id()
))
}
#[test]
fn batch_collection_is_recursive_and_sorted() {
let root = temporary_directory();
let nested = root.join("nested");
std::fs::create_dir_all(&nested).unwrap();
std::fs::write(root.join("b.wav"), []).unwrap();
std::fs::write(root.join("ignore.txt"), []).unwrap();
std::fs::write(nested.join("a.FLAC"), []).unwrap();
assert_eq!(
collect_batch_files(&root, false).unwrap(),
vec![root.join("b.wav")]
);
assert_eq!(
collect_batch_files(&root, true).unwrap(),
vec![root.join("b.wav"), nested.join("a.FLAC")]
);
std::fs::remove_dir_all(root).unwrap();
}
#[test]
fn validates_batch_output_format() {
assert_eq!(normalize_output_extension(".flac").unwrap(), "flac");
assert_eq!(normalize_output_extension("aac").unwrap(), "aac");
assert!(normalize_output_extension("wma").is_err());
}
#[test]
fn batch_counts_distinguish_completed_skipped_and_failed_results() {
let results = [
Ok(BatchFileOutcome::Completed),
Ok(BatchFileOutcome::Skipped),
Err("processing failed"),
Err("cancelled"),
];
assert_eq!(
count_batch_results(&results),
BatchCounts {
succeeded: 1,
skipped: 1,
failed: 2,
}
);
}
#[test]
fn batch_processes_nested_audio_and_converts_format() {
let root = temporary_directory();
let input = root.join("input");
let output = root.join("output");
std::fs::create_dir_all(input.join("nested")).unwrap();
let audio = denoize::Audio {
sample_rate: 16_000,
channels: vec![vec![0.0; 3_200]],
bits_per_sample: 16,
sample_format: hound::SampleFormat::Int,
channel_mask: None,
};
denoize::write_audio(
input.join("nested/sample.wav"),
&audio,
EncodeOptions::default(),
)
.unwrap();
let options = Overrides {
batch: true,
recursive: true,
jobs: Some(2),
output_format: Some("flac".into()),
..Overrides::default()
};
run_batch(input.to_str().unwrap(), output.to_str().unwrap(), &options).unwrap();
assert!(output.join("nested/sample.flac").is_file());
std::fs::remove_dir_all(root).unwrap();
}
#[test]
fn deterministic_batch_is_byte_stable_even_with_multiple_requested_jobs() {
let root = temporary_directory();
let input = root.join("input");
let output = root.join("output");
std::fs::create_dir_all(&input).unwrap();
for (name, frequency) in [("a.wav", 220.0), ("b.wav", 440.0)] {
let audio = denoize::Audio {
sample_rate: 16_000,
channels: vec![(0..3_200)
.map(|index| {
(2.0 * std::f64::consts::PI * frequency * index as f64 / 16_000.0)
.sin()
* 0.2
})
.collect()],
bits_per_sample: 16,
sample_format: hound::SampleFormat::Int,
channel_mask: None,
};
denoize::write_audio(input.join(name), &audio, EncodeOptions::default()).unwrap();
}
let options = Overrides {
batch: true,
deterministic: true,
force: true,
jobs: Some(8),
no_progress: true,
..Overrides::default()
};
run_batch(input.to_str().unwrap(), output.to_str().unwrap(), &options).unwrap();
let first_a = std::fs::read(output.join("a.wav")).unwrap();
let first_b = std::fs::read(output.join("b.wav")).unwrap();
run_batch(input.to_str().unwrap(), output.to_str().unwrap(), &options).unwrap();
assert_eq!(first_a, std::fs::read(output.join("a.wav")).unwrap());
assert_eq!(first_b, std::fs::read(output.join("b.wav")).unwrap());
std::fs::remove_dir_all(root).unwrap();
}
#[test]
fn resume_skips_outputs_recorded_as_complete() {
let root = temporary_directory();
let input = root.join("input");
let output = root.join("output");
std::fs::create_dir_all(&input).unwrap();
let audio = denoize::Audio {
sample_rate: 16_000,
channels: vec![vec![0.0; 1_600]],
bits_per_sample: 16,
sample_format: hound::SampleFormat::Int,
channel_mask: None,
};
denoize::write_audio(input.join("sample.wav"), &audio, EncodeOptions::default()).unwrap();
let options = Overrides {
batch: true,
resume: true,
no_progress: true,
..Overrides::default()
};
run_batch(input.to_str().unwrap(), output.to_str().unwrap(), &options).unwrap();
let first_modified = std::fs::metadata(output.join("sample.wav"))
.unwrap()
.modified()
.unwrap();
run_batch(input.to_str().unwrap(), output.to_str().unwrap(), &options).unwrap();
let second_modified = std::fs::metadata(output.join("sample.wav"))
.unwrap()
.modified()
.unwrap();
assert_eq!(first_modified, second_modified);
assert!(read_batch_state(&output.join(".denoize-state"))
.unwrap()
.contains("sample.wav"));
std::fs::remove_dir_all(root).unwrap();
}
}
#[cfg(test)]
mod auto_backend_tests {
use super::*;
#[test]
fn parses_auto_backend() {
let (_, _, options) = parse_args(&[
"input.wav".into(),
"output.wav".into(),
"--backend".into(),
"auto".into(),
])
.unwrap();
assert!(options.auto_backend);
assert!(options.backend.is_none());
}
#[test]
fn automatic_selection_uses_an_available_backend() {
let selected = service::select_backend(BackendChoice::Auto, 30.0, None);
assert!(Backend::available_names().contains(&service::backend_name(selected)));
}
}
#[cfg(test)]
mod streaming_tests {
use super::*;
#[test]
fn parses_stream_option() {
let (_, _, options) = parse_args(&[
"input.wav".into(),
"output.wav".into(),
"--stream".into(),
"--stream-frames".into(),
"4096".into(),
"--max-memory".into(),
"64".into(),
])
.unwrap();
assert!(options.stream);
assert_eq!(options.stream_frames, Some(4096));
assert_eq!(options.max_memory_mb, Some(64));
}
#[test]
fn rejects_zero_resource_limits() {
let error = validate_resource_options(&Overrides {
max_memory_mb: Some(0),
..Overrides::default()
})
.unwrap_err();
assert!(error.contains("--max-memory"));
let error = validate_resource_options(&Overrides {
stream_frames: Some(0),
..Overrides::default()
})
.unwrap_err();
assert!(error.contains("--stream-frames"));
}
#[test]
fn streams_wav_without_loading_the_complete_audio() {
let root = std::env::temp_dir().join(format!(
"denoize-stream-test-{}-{}",
std::process::id(),
std::thread::current().name().unwrap_or("test")
));
std::fs::create_dir_all(&root).unwrap();
let input = root.join("input.wav");
let output = root.join("output.wav");
let spec = hound::WavSpec {
channels: 1,
sample_rate: 16_000,
bits_per_sample: 16,
sample_format: hound::SampleFormat::Int,
};
let mut writer = hound::WavWriter::create(&input, spec).unwrap();
for frame in 0..20_000 {
let sample = (0.2
* (2.0 * std::f64::consts::PI * 440.0 * frame as f64 / spec.sample_rate as f64)
.sin()
* 32_767.0) as i16;
writer.write_sample(sample).unwrap();
}
writer.finalize().unwrap();
run_streaming_wav(
input.to_str().unwrap(),
output.to_str().unwrap(),
Overrides {
stream: true,
stream_frames: Some(257),
..Overrides::default()
},
)
.unwrap();
let result = read_audio(&output).unwrap();
assert_eq!(result.sample_rate, spec.sample_rate);
assert_eq!(result.channels(), 1);
assert_eq!(result.frames(), 20_000);
std::fs::remove_dir_all(root).unwrap();
}
}
#[cfg(test)]
mod config_file_tests {
use super::*;
#[test]
fn parses_toml_defaults() {
let options = parse_config(
r#"
backend = "auto"
preset = "hifi"
mode = "speech"
strength = 0.42
adaptive_noise = true
vad = true
preserve_metadata = false
downmix = "stereo"
deterministic = true
seed = 12345
stream_frames = 4096
max_memory_mb = 64
"#,
"test.toml",
)
.unwrap();
assert!(options.auto_backend);
assert!(options.deterministic);
assert_eq!(options.seed, Some(12345));
assert_eq!(options.downmix, Some(DownmixMode::Stereo));
assert_eq!(options.preset, Some(Preset::HiFi));
assert_eq!(options.mode, Some(ProcessingMode::Speech));
assert_eq!(options.strength, Some(0.42));
assert!(options.adaptive_noise && options.vad && options.no_metadata);
assert_eq!(options.stream_frames, Some(4096));
assert_eq!(options.max_memory_mb, Some(64));
}
#[test]
fn parses_desktop_exported_config() {
let options = parse_config(
r#"
backend = "auto"
preset = "hifi"
mode = "speech"
strength = 0.42
adaptive_noise = true
vad = true
channels = "linked"
downmix = "stereo"
loudness_lufs = -16.0
true_peak_dbtp = -1.0
preserve_metadata = false
force = true
mp3_bitrate_kbps = 256
m4a_bitrate_kbps = 224
aac_encoder = "oxide"
onnx_model = "model.onnx"
onnx_rate = 48000
sgmse_profile = "quality"
deterministic = true
"#,
"desktop.toml",
)
.unwrap();
assert!(options.auto_backend);
assert_eq!(options.preset, Some(Preset::HiFi));
assert_eq!(options.mode, Some(ProcessingMode::Speech));
assert_eq!(options.strength, Some(0.42));
assert!(options.adaptive_noise);
assert!(options.vad);
assert_eq!(options.channel_mode, Some(ChannelMode::StereoLinked));
assert_eq!(options.downmix, Some(DownmixMode::Stereo));
assert_eq!(options.loudness_lufs, Some(-16.0));
assert_eq!(options.true_peak_dbtp, Some(-1.0));
assert!(options.no_metadata);
assert!(options.force);
assert_eq!(options.mp3_bitrate_kbps, Some(256));
assert_eq!(options.m4a_bitrate_kbps, Some(224));
assert_eq!(options.aac_encoder, Some(AacEncoder::Oxide));
assert_eq!(options.onnx_model.as_deref(), Some("model.onnx"));
assert_eq!(options.onnx_sample_rate, Some(48_000));
assert_eq!(options.sgmse_profile, Some(SgmseProfile::Quality));
assert!(options.deterministic);
}
#[test]
fn rejects_invalid_desktop_enum_values() {
let error = parse_config("aac_encoder = \"invalid\"", "desktop.toml").unwrap_err();
assert!(error.contains("unknown AAC encoder in config: invalid"));
let error = parse_config("sgmse_profile = \"invalid\"", "desktop.toml").unwrap_err();
assert!(error.contains("unknown SGMSE profile in config: invalid"));
}
#[test]
fn accepts_legacy_desktop_true_peak_without_loudness() {
let options = parse_config(
"true_peak_dbtp = -1.0\nmp3_bitrate_kbps = 192\n",
"legacy-desktop.toml",
)
.unwrap();
assert_eq!(options.loudness_lufs, None);
assert_eq!(options.true_peak_dbtp, None);
let explicit = parse_config("true_peak_dbtp = -0.5", "manual.toml").unwrap();
assert_eq!(explicit.true_peak_dbtp, Some(-0.5));
}
#[test]
fn rejects_unknown_config_keys() {
let error = parse_config("strenth = 0.5", "test.toml").unwrap_err();
assert!(error.contains("unknown field"));
}
#[test]
fn command_line_overrides_config_defaults() {
let path = std::env::temp_dir().join(format!(
"denoize-config-{}-{}.toml",
std::process::id(),
std::thread::current().name().unwrap_or("test")
));
std::fs::write(&path, "backend = \"auto\"\nstrength = 0.25\n").unwrap();
let args = vec![
"input.wav".into(),
"output.wav".into(),
"--config".into(),
path.to_string_lossy().into_owned(),
"--backend".into(),
"classical".into(),
"--strength".into(),
"0.75".into(),
];
let (_, _, options) = parse_args(&args).unwrap();
std::fs::remove_file(path).unwrap();
assert_eq!(options.backend, Some(Backend::Classical));
assert!(!options.auto_backend);
assert_eq!(options.strength, Some(0.75));
}
#[test]
fn parses_explicit_downmix_mode() {
let (_, _, options) = parse_args(&[
"input.wav".into(),
"output.mp3".into(),
"--downmix".into(),
"stereo".into(),
])
.unwrap();
assert_eq!(options.downmix, Some(DownmixMode::Stereo));
}
#[test]
fn parses_deterministic_seed_and_implies_mode() {
let (_, _, options) = parse_args(&[
"input.wav".into(),
"output.wav".into(),
"--seed".into(),
"42".into(),
])
.unwrap();
assert!(options.deterministic);
assert_eq!(options.seed, Some(42));
}
}
fn run_metrics(args: &[String]) -> Result<(), String> {
let reference = args.first().ok_or("metrics requires REFERENCE and TEST")?;
let test = args.get(1).ok_or("metrics requires REFERENCE and TEST")?;
let report =
denoize::benchmark::BenchmarkReport::compare(&read_audio(reference)?, &read_audio(test)?)?;
if args.iter().any(|argument| argument == "--json") {
println!("{}", report.json());
} else {
println!("{}", report.markdown());
}
Ok(())
}
fn run_compare(args: &[String]) -> Result<(), String> {
if args.len() < 3 {
return Err("compare requires CLEAN NOISY ENHANCED".into());
}
if args[3..]
.iter()
.any(|argument| argument != "--json" && argument != "--html")
{
return Err("compare accepts only --json or --html after the input files".into());
}
if args.iter().any(|argument| argument == "--json")
&& args.iter().any(|argument| argument == "--html")
{
return Err("compare accepts only one output format".into());
}
let clean = args
.first()
.ok_or("compare requires CLEAN NOISY ENHANCED")?;
let noisy = args.get(1).ok_or("compare requires CLEAN NOISY ENHANCED")?;
let enhanced = args.get(2).ok_or("compare requires CLEAN NOISY ENHANCED")?;
let report = denoize::benchmark::ComparisonReport::compare(
&read_audio(clean)?,
&read_audio(noisy)?,
&read_audio(enhanced)?,
)?;
if args.iter().any(|argument| argument == "--json") {
println!("{}", report.json());
} else if args.iter().any(|argument| argument == "--html") {
println!("{}", report.html());
} else {
println!("{}", report.markdown());
}
Ok(())
}
fn models_usage() -> &'static str {
"\
Manage verified external models.
USAGE:
denoize models list
denoize models info <MODEL|all>
denoize models install <MODEL|all> [DOWNLOAD OPTIONS]
denoize models install <MODEL> --from <PATH>
denoize models update <MODEL|all> [DOWNLOAD OPTIONS]
denoize models verify <MODEL|all>
denoize models remove <MODEL|all>
denoize models path <MODEL|all>
denoize models cache-dir
DOWNLOAD OPTIONS:
--offline never access the network; use only verified cached data
--proxy <URL> use this proxy instead of proxy environment variables
--no-proxy connect directly and ignore proxy environment variables
--url <URL> download one MODEL from an alternate HTTP(S) URL
--bearer-token-env <VAR> read a bearer token from environment variable VAR
--basic-user <USER> username for HTTP Basic authentication
--basic-password-env <VAR> read the Basic password from environment variable VAR
--from <PATH> install one MODEL from a local file (install only)
Bearer tokens and Basic passwords are read from environment variables instead
of literal secret flags. Basic authentication requires both --basic-user and
--basic-password-env. Signed --url values and proxy credentials can still be
visible in process arguments. Alternate sources, origin authentication, and
--from accept one model, not `all`; --url rejects userinfo credentials.
ENVIRONMENT:
DENOIZE_MODEL_OFFLINE, DENOIZE_MODEL_URL, DENOIZE_MODEL_PROXY,
DENOIZE_MODEL_BEARER_TOKEN, DENOIZE_MODEL_USERNAME, DENOIZE_MODEL_PASSWORD
HTTPS_PROXY, HTTP_PROXY, ALL_PROXY, NO_PROXY (and lowercase variants)
"
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum ModelCommand {
Info,
Install,
Update,
Verify,
Remove,
Path,
}
#[derive(Debug)]
enum ParsedModelsCommand {
Help,
List,
CacheDir,
Run {
command: ModelCommand,
target: String,
download_options: Option<Box<denoize::models::ModelDownloadOptions>>,
source_file: Option<std::path::PathBuf>,
},
}
fn models_option_value(args: &[String], index: &mut usize, flag: &str) -> Result<String, String> {
*index += 1;
let value = args
.get(*index)
.ok_or_else(|| format!("missing value for {flag}"))?;
if value.is_empty() {
return Err(format!("empty value for {flag}"));
}
Ok(value.clone())
}
fn validate_model_source_url(value: &str) -> Result<(), String> {
let source = url::Url::parse(value)
.map_err(|_| "invalid value for --url: expected an HTTP(S) URL".to_string())?;
if !matches!(source.scheme(), "http" | "https") {
return Err("invalid value for --url: expected an HTTP(S) URL".into());
}
if !source.username().is_empty() || source.password().is_some() {
return Err(
"--url must not contain credentials; use --bearer-token-env or Basic authentication options"
.into(),
);
}
Ok(())
}
fn read_model_secret<F>(
flag: &str,
variable: &str,
read_environment: &mut F,
) -> Result<String, String>
where
F: FnMut(&str) -> Result<String, String>,
{
if variable.trim().is_empty() {
return Err(format!("empty environment variable name for {flag}"));
}
let secret = read_environment(variable).map_err(|error| {
format!("failed to read environment variable {variable} for {flag}: {error}")
})?;
if secret.is_empty() {
return Err(format!(
"environment variable {variable} referenced by {flag} is empty"
));
}
Ok(secret)
}
fn parse_models_command<F>(
args: &[String],
mut download_options: denoize::models::ModelDownloadOptions,
mut read_environment: F,
) -> Result<ParsedModelsCommand, String>
where
F: FnMut(&str) -> Result<String, String>,
{
if args
.iter()
.any(|argument| matches!(argument.as_str(), "-h" | "--help"))
|| args.first().map(String::as_str) == Some("help")
{
return Ok(ParsedModelsCommand::Help);
}
let command_name = args.first().map(String::as_str).unwrap_or("list");
if matches!(command_name, "list" | "cache-dir") {
if args.len() > 1 {
return Err(format!("models {command_name} accepts no arguments"));
}
return Ok(if command_name == "list" {
ParsedModelsCommand::List
} else {
ParsedModelsCommand::CacheDir
});
}
let command = match command_name {
"info" => ModelCommand::Info,
"install" => ModelCommand::Install,
"update" => ModelCommand::Update,
"verify" => ModelCommand::Verify,
"remove" => ModelCommand::Remove,
"path" => ModelCommand::Path,
_ => return Err(format!("unknown models command: {command_name}")),
};
let target = args
.get(1)
.filter(|target| !target.starts_with('-'))
.ok_or_else(|| format!("models {command_name} requires MODEL|all"))?
.clone();
if !matches!(command, ModelCommand::Install | ModelCommand::Update) {
if args.len() > 2 {
return Err(format!(
"models {command_name} does not accept options or extra arguments"
));
}
return Ok(ParsedModelsCommand::Run {
command,
target,
download_options: None,
source_file: None,
});
}
let mut offline_seen = false;
let mut proxy_flag: Option<&str> = None;
let mut source_url_seen = false;
let mut bearer_variable: Option<String> = None;
let mut basic_user: Option<String> = None;
let mut basic_password_variable: Option<String> = None;
let mut source_file: Option<std::path::PathBuf> = None;
let mut index = 2;
while index < args.len() {
let flag = args[index].as_str();
match flag {
"--offline" => {
if offline_seen {
return Err("--offline specified more than once".into());
}
offline_seen = true;
download_options.offline = true;
}
"--proxy" => {
if let Some(previous) = proxy_flag {
return Err(format!("--proxy cannot be combined with {previous}"));
}
let value = models_option_value(args, &mut index, flag)?;
proxy_flag = Some("--proxy");
download_options.proxy = denoize::models::ModelProxy::Url(value);
}
"--no-proxy" => {
if let Some(previous) = proxy_flag {
return Err(format!("--no-proxy cannot be combined with {previous}"));
}
proxy_flag = Some("--no-proxy");
download_options.proxy = denoize::models::ModelProxy::Disabled;
}
"--url" => {
if source_url_seen {
return Err("--url specified more than once".into());
}
let value = models_option_value(args, &mut index, flag)?;
validate_model_source_url(&value)?;
source_url_seen = true;
download_options.source_url = Some(value);
}
"--bearer-token-env" => {
if bearer_variable.is_some() {
return Err("--bearer-token-env specified more than once".into());
}
bearer_variable = Some(models_option_value(args, &mut index, flag)?);
}
"--basic-user" => {
if basic_user.is_some() {
return Err("--basic-user specified more than once".into());
}
basic_user = Some(models_option_value(args, &mut index, flag)?);
}
"--basic-password-env" => {
if basic_password_variable.is_some() {
return Err("--basic-password-env specified more than once".into());
}
basic_password_variable = Some(models_option_value(args, &mut index, flag)?);
}
"--from" => {
if source_file.is_some() {
return Err("--from specified more than once".into());
}
source_file = Some(models_option_value(args, &mut index, flag)?.into());
}
value if value.starts_with('-') => {
return Err(format!("unknown models {command_name} option: {value}"));
}
value => {
return Err(format!(
"unexpected argument for models {command_name}: {value}"
));
}
}
index += 1;
}
if source_file.is_some() {
if command != ModelCommand::Install {
return Err("--from is supported only by `models install`".into());
}
if target == "all" {
return Err("--from requires one MODEL and cannot be used with `all`".into());
}
if source_url_seen
|| proxy_flag.is_some()
|| bearer_variable.is_some()
|| basic_user.is_some()
|| basic_password_variable.is_some()
{
return Err("--from cannot be combined with network download options".into());
}
download_options = denoize::models::ModelDownloadOptions::default();
download_options.offline = offline_seen;
}
if bearer_variable.is_some() && (basic_user.is_some() || basic_password_variable.is_some()) {
return Err(
"--bearer-token-env cannot be combined with Basic authentication options".into(),
);
}
download_options.authentication = if let Some(variable) = bearer_variable {
Some(denoize::models::ModelAuthentication::Bearer(
read_model_secret("--bearer-token-env", &variable, &mut read_environment)?,
))
} else {
match (basic_user, basic_password_variable) {
(Some(username), Some(variable)) => {
let password =
read_model_secret("--basic-password-env", &variable, &mut read_environment)?;
Some(denoize::models::ModelAuthentication::Basic { username, password })
}
(None, None) => download_options.authentication,
_ => {
return Err(
"--basic-user and --basic-password-env must be specified together".into(),
)
}
}
};
if target == "all" && download_options.source_url.is_some() {
return Err(
"an alternate model URL requires one MODEL and cannot be used with `all`".into(),
);
}
if target == "all" && download_options.authentication.is_some() {
return Err("model authentication requires one MODEL and cannot be used with `all`".into());
}
Ok(ParsedModelsCommand::Run {
command,
target,
download_options: Some(Box::new(download_options)),
source_file,
})
}
fn model_download_options_from_environment_with<F>(
args: &[String],
mut read_environment: F,
) -> Result<denoize::models::ModelDownloadOptions, String>
where
F: FnMut(&str) -> Option<String>,
{
if args.iter().any(|argument| argument == "--from") {
return Ok(denoize::models::ModelDownloadOptions::default());
}
let overrides_offline = args.iter().any(|argument| argument == "--offline");
let overrides_source = args.iter().any(|argument| argument == "--url");
let overrides_proxy = args
.iter()
.any(|argument| matches!(argument.as_str(), "--proxy" | "--no-proxy"));
let overrides_authentication = args.iter().any(|argument| {
matches!(
argument.as_str(),
"--bearer-token-env" | "--basic-user" | "--basic-password-env"
)
});
denoize::models::ModelDownloadOptions::from_env_with(|name| {
let overridden = match name {
"DENOIZE_MODEL_OFFLINE" => overrides_offline,
"DENOIZE_MODEL_URL" => overrides_source,
"DENOIZE_MODEL_PROXY" => overrides_proxy,
"DENOIZE_MODEL_BEARER_TOKEN" | "DENOIZE_MODEL_USERNAME" | "DENOIZE_MODEL_PASSWORD" => {
overrides_authentication
}
_ => false,
};
(!overridden).then(|| read_environment(name)).flatten()
})
}
fn run_models(args: &[String]) -> Result<(), String> {
let help_requested = args
.iter()
.any(|argument| matches!(argument.as_str(), "-h" | "--help"))
|| args.first().map(String::as_str) == Some("help");
let download_command = matches!(args.first().map(String::as_str), Some("install" | "update"));
let download_options = if download_command && !help_requested {
model_download_options_from_environment_with(args, |name| std::env::var(name).ok())?
} else {
denoize::models::ModelDownloadOptions::default()
};
let parsed = parse_models_command(args, download_options, |name| {
std::env::var(name).map_err(|error| error.to_string())
})?;
let (command, target, download_options, source_file) = match parsed {
ParsedModelsCommand::Help => {
print!("{}", models_usage());
return Ok(());
}
ParsedModelsCommand::List => {
println!("NAME\tBACKEND\tRATE\tLICENSE\tSTATUS");
for model in denoize::models::MODELS {
let status = if denoize::models::verify(model).is_ok() {
"installed"
} else {
"not-installed"
};
println!(
"{}\t{}\t{}\t{}\t{}",
model.name, model.backend, model.sample_rate, model.license, status
);
}
return Ok(());
}
ParsedModelsCommand::CacheDir => {
println!("{}", denoize::models::cache_dir()?.display());
return Ok(());
}
ParsedModelsCommand::Run {
command,
target,
download_options,
source_file,
} => (command, target, download_options, source_file),
};
let models: Vec<_> = if target == "all" {
denoize::models::MODELS.iter().collect()
} else {
vec![denoize::models::find(&target)
.ok_or_else(|| format!("unknown model: {target} (run `denoize models list`)"))?]
};
for model in models {
match command {
ModelCommand::Info => {
println!("name: {}", model.name);
println!("backend: {}", model.backend);
println!("sample-rate: {}", model.sample_rate);
println!("license: {}", model.license);
println!("revision: {}", model.revision);
println!("sha256: {}", model.sha256);
println!("url: {}", denoize::models::redact_url(model.url));
println!("path: {}", denoize::models::path(model)?.display());
}
ModelCommand::Install => {
let installed = if let Some(source) = source_file.as_ref() {
denoize::models::install_from_file(model, source)?
} else {
denoize::models::install_with_options(
model,
download_options
.as_ref()
.expect("download options exist for install"),
)?
};
println!("{}", installed.display());
}
ModelCommand::Update => println!(
"{}",
denoize::models::update_with_options(
model,
download_options
.as_ref()
.expect("download options exist for update"),
)?
.display()
),
ModelCommand::Verify => {
println!("verified {}", denoize::models::verify(model)?.display())
}
ModelCommand::Remove => println!(
"{} {}",
if denoize::models::remove(model)? {
"removed"
} else {
"not-installed"
},
model.name
),
ModelCommand::Path => println!("{}", denoize::models::path(model)?.display()),
}
}
Ok(())
}
#[cfg(test)]
mod model_command_tests {
use super::*;
fn missing_secret(name: &str) -> Result<String, String> {
Err(format!("{name} is not set"))
}
#[test]
fn explicit_model_flags_override_invalid_environment_defaults() {
let args = vec![
"install".into(),
"gtcrn-dns3".into(),
"--offline".into(),
"--url".into(),
"https://models.example/model.onnx".into(),
"--no-proxy".into(),
"--bearer-token-env".into(),
"MODEL_TOKEN".into(),
];
let options = model_download_options_from_environment_with(&args, |name| {
Some(
match name {
"DENOIZE_MODEL_OFFLINE" => "not-a-boolean",
"DENOIZE_MODEL_URL" => "environment-url",
"DENOIZE_MODEL_PROXY" => "environment-proxy",
"DENOIZE_MODEL_BEARER_TOKEN" => "environment-bearer",
"DENOIZE_MODEL_USERNAME" => "environment-user",
"DENOIZE_MODEL_PASSWORD" => "environment-password",
_ => return None,
}
.into(),
)
})
.unwrap();
assert!(!options.offline);
assert!(options.source_url.is_none());
assert!(matches!(
options.proxy,
denoize::models::ModelProxy::Environment
));
assert!(options.authentication.is_none());
}
#[test]
fn local_model_install_does_not_validate_unrelated_environment_defaults() {
let args = vec![
"install".into(),
"gtcrn-dns3".into(),
"--from".into(),
"model.onnx".into(),
];
let options = model_download_options_from_environment_with(&args, |_| {
panic!("local installs must not read model download environment variables")
})
.unwrap();
assert!(!options.offline);
assert!(options.source_url.is_none());
assert!(options.authentication.is_none());
}
#[test]
fn parses_model_download_overrides_without_reading_process_environment() {
let mut base = denoize::models::ModelDownloadOptions::default();
base.source_url = Some("https://environment.invalid/model".into());
base.authentication = Some(denoize::models::ModelAuthentication::Basic {
username: "environment-user".into(),
password: "environment-secret".into(),
});
let args = vec![
"update".into(),
"gtcrn-dns3".into(),
"--url".into(),
"https://models.example/model.onnx".into(),
"--no-proxy".into(),
"--bearer-token-env".into(),
"MODEL_TOKEN".into(),
];
let parsed = parse_models_command(&args, base, |name| {
assert_eq!(name, "MODEL_TOKEN");
Ok("secret-token".into())
})
.unwrap();
let ParsedModelsCommand::Run {
command,
target,
download_options: Some(options),
source_file,
} = parsed
else {
panic!("expected an executable model command");
};
assert_eq!(command, ModelCommand::Update);
assert_eq!(target, "gtcrn-dns3");
assert!(source_file.is_none());
assert_eq!(
options.source_url.as_deref(),
Some("https://models.example/model.onnx")
);
assert!(matches!(
options.proxy,
denoize::models::ModelProxy::Disabled
));
assert!(matches!(
options.authentication,
Some(denoize::models::ModelAuthentication::Bearer(ref token)) if token == "secret-token"
));
}
#[test]
fn parses_basic_authentication_and_local_install() {
let basic = vec![
"install".into(),
"gtcrn-dns3".into(),
"--basic-user".into(),
"release-bot".into(),
"--basic-password-env".into(),
"MODEL_PASSWORD".into(),
];
let parsed = parse_models_command(
&basic,
denoize::models::ModelDownloadOptions::default(),
|_| Ok("password-from-environment".into()),
)
.unwrap();
let ParsedModelsCommand::Run {
download_options: Some(options),
..
} = parsed
else {
panic!("expected download options");
};
assert!(matches!(
options.authentication,
Some(denoize::models::ModelAuthentication::Basic {
ref username,
ref password,
}) if username == "release-bot" && password == "password-from-environment"
));
let local = vec![
"install".into(),
"gtcrn-dns3".into(),
"--offline".into(),
"--from".into(),
"model.onnx".into(),
];
let parsed = parse_models_command(
&local,
denoize::models::ModelDownloadOptions::default(),
missing_secret,
)
.unwrap();
let ParsedModelsCommand::Run {
command,
source_file: Some(source),
download_options: Some(options),
..
} = parsed
else {
panic!("expected a local install");
};
assert_eq!(command, ModelCommand::Install);
assert_eq!(source, std::path::PathBuf::from("model.onnx"));
assert!(options.offline);
}
#[test]
fn rejects_conflicting_or_incomplete_model_options() {
let cases = [
(
vec![
"install".into(),
"gtcrn-dns3".into(),
"--proxy".into(),
"http://proxy.example".into(),
"--no-proxy".into(),
],
"cannot be combined",
),
(
vec![
"install".into(),
"gtcrn-dns3".into(),
"--basic-user".into(),
"release-bot".into(),
],
"must be specified together",
),
(
vec![
"install".into(),
"gtcrn-dns3".into(),
"--bearer-token-env".into(),
"TOKEN".into(),
"--basic-user".into(),
"release-bot".into(),
"--basic-password-env".into(),
"PASSWORD".into(),
],
"cannot be combined",
),
(
vec![
"install".into(),
"gtcrn-dns3".into(),
"--from".into(),
"model.onnx".into(),
"--proxy".into(),
"http://proxy.example".into(),
],
"network download options",
),
];
for (args, expected) in cases {
let error = parse_models_command(
&args,
denoize::models::ModelDownloadOptions::default(),
missing_secret,
)
.unwrap_err();
assert!(error.contains(expected), "unexpected error: {error}");
}
}
#[test]
fn rejects_options_outside_their_supported_target_or_command() {
let cases = [
(
vec!["info".into(), "gtcrn-dns3".into(), "--offline".into()],
"does not accept options",
),
(
vec![
"update".into(),
"gtcrn-dns3".into(),
"--from".into(),
"model.onnx".into(),
],
"install",
),
(
vec![
"install".into(),
"all".into(),
"--from".into(),
"model.onnx".into(),
],
"cannot be used with `all`",
),
(
vec![
"update".into(),
"all".into(),
"--url".into(),
"https://models.example/model.onnx".into(),
],
"cannot be used with `all`",
),
(
vec![
"install".into(),
"gtcrn-dns3".into(),
"--url".into(),
"https://user:secret@models.example/model.onnx".into(),
],
"must not contain credentials",
),
];
for (args, expected) in cases {
let error = parse_models_command(
&args,
denoize::models::ModelDownloadOptions::default(),
missing_secret,
)
.unwrap_err();
assert!(error.contains(expected), "unexpected error: {error}");
}
}
#[test]
fn rejects_environment_source_or_authentication_for_all_models() {
let args = vec!["update".into(), "all".into()];
let mut source = denoize::models::ModelDownloadOptions::default();
source.source_url = Some("https://mirror.example/model.onnx".into());
let source_error = parse_models_command(&args, source, missing_secret).unwrap_err();
assert!(source_error.contains("cannot be used with `all`"));
let mut authenticated = denoize::models::ModelDownloadOptions::default();
authenticated.authentication = Some(denoize::models::ModelAuthentication::Bearer(
"environment-token".into(),
));
let authentication_error =
parse_models_command(&args, authenticated, missing_secret).unwrap_err();
assert!(authentication_error.contains("requires one MODEL"));
}
#[test]
fn reports_missing_secret_environment_variables_without_exposing_values() {
let args = vec![
"install".into(),
"gtcrn-dns3".into(),
"--bearer-token-env".into(),
"MISSING_TOKEN".into(),
];
let error = parse_models_command(
&args,
denoize::models::ModelDownloadOptions::default(),
missing_secret,
)
.unwrap_err();
assert!(error.contains("MISSING_TOKEN"));
assert!(error.contains("not set"));
}
#[test]
fn exposes_dedicated_models_help() {
let parsed = parse_models_command(
&["--help".into()],
denoize::models::ModelDownloadOptions::default(),
missing_secret,
)
.unwrap();
assert!(matches!(parsed, ParsedModelsCommand::Help));
for flag in [
"--offline",
"--proxy",
"--no-proxy",
"--url",
"--bearer-token-env",
"--basic-user",
"--basic-password-env",
"--from",
] {
assert!(models_usage().contains(flag));
}
}
}
fn main() {
let args: Vec<String> = std::env::args().skip(1).collect();
if let Err(e) = run(&args) {
eprintln!("denoize: error: {e}");
eprintln!("run 'denoize --help' for usage.");
std::process::exit(1);
}
}
#[cfg(all(test, feature = "onnx"))]
mod tests {
use super::*;
#[test]
fn parses_onnx_model_options() {
let args = vec![
"input.wav".into(),
"output.wav".into(),
"--backend".into(),
"onnx".into(),
"--onnx-model".into(),
"model.onnx".into(),
"--onnx-rate".into(),
"48000".into(),
];
let (_, _, options) = parse_args(&args).unwrap();
assert_eq!(options.backend, Some(Backend::Onnx));
assert_eq!(options.onnx_model.as_deref(), Some("model.onnx"));
assert_eq!(options.onnx_sample_rate, Some(48_000));
}
#[test]
fn parses_live_device_options() {
let args = vec![
"-".into(),
"-".into(),
"--input-device".into(),
"Mic".into(),
"--output-device".into(),
"Cable".into(),
"--chunk-ms".into(),
"40".into(),
];
let (_, _, options) = parse_args(&args).unwrap();
assert_eq!(options.input_device.as_deref(), Some("Mic"));
assert_eq!(options.output_device.as_deref(), Some("Cable"));
assert_eq!(options.chunk_ms, Some(40));
}
}