use crate::engine::stage::{BLOCK_FRAMES, BlockBuf, Stage, StageCtx};
use crate::engine::stages::band_split::{KEYLOCK_CROSSOVER_HZ, TwoBandSplit};
use crate::engine::stages::delay::FixedDelay;
use crate::engine::stages::sola::SolaCorrector;
pub(crate) const KEYLOCK_LATENCY_FRAMES: usize = 560;
pub const CORRECTION_FADE_START_DEV: f64 = 0.205;
pub const CORRECTION_FADE_END_DEV: f64 = 0.35;
pub const KEYLOCK_TOGGLE_FADE_FRAMES: usize = 512;
#[derive(Debug)]
pub(crate) struct KeylockStage {
split: TwoBandSplit,
low_delay: FixedDelay,
raw_high_delay: FixedDelay,
sola: SolaCorrector,
low: Vec<[f32; BLOCK_FRAMES]>,
high: Vec<[f32; BLOCK_FRAMES]>,
high_raw: Vec<[f32; BLOCK_FRAMES]>,
enable: f32,
}
impl KeylockStage {
pub(crate) fn new(sample_rate: u32, channels: usize) -> Self {
let sola = SolaCorrector::new(channels, KEYLOCK_LATENCY_FRAMES);
debug_assert_eq!(sola.latency_frames(), KEYLOCK_LATENCY_FRAMES);
Self {
split: TwoBandSplit::new(KEYLOCK_CROSSOVER_HZ, sample_rate, channels),
low_delay: FixedDelay::new(KEYLOCK_LATENCY_FRAMES, channels),
raw_high_delay: FixedDelay::new(KEYLOCK_LATENCY_FRAMES, channels),
sola,
low: vec![[0.0; BLOCK_FRAMES]; channels],
high: vec![[0.0; BLOCK_FRAMES]; channels],
high_raw: vec![[0.0; BLOCK_FRAMES]; channels],
enable: f32::NAN,
}
}
}
impl Stage for KeylockStage {
fn process(&mut self, block: &mut BlockBuf, ctx: &StageCtx<'_>) {
let transposition = if ctx.embedded_rate.is_finite() && ctx.embedded_rate > 0.0 {
1.0 / ctx.embedded_rate
} else {
1.0
};
self.sola.set_transposition(transposition);
self.sola.set_rate_slope(ctx.embedded_rate_slope);
for ch in 0..block.channels() {
let (low, high) = (&mut self.low[ch], &mut self.high[ch]);
self.split.process_channel(ch, block.channel(ch), low, high);
self.high_raw[ch].copy_from_slice(high);
self.raw_high_delay
.process_channel(ch, &mut self.high_raw[ch]);
self.low_delay.process_channel(ch, low);
}
self.sola.process_block(&mut self.high, ctx.onsets);
let deviation = (ctx.embedded_rate - 1.0).abs();
let correction = ((CORRECTION_FADE_END_DEV - deviation)
/ (CORRECTION_FADE_END_DEV - CORRECTION_FADE_START_DEV))
.clamp(0.0, 1.0) as f32;
let target = (ctx.keylock.clamp(0.0, 1.0)) as f32;
if self.enable.is_nan() {
self.enable = target;
}
let step = 1.0 / KEYLOCK_TOGGLE_FADE_FRAMES as f32;
let mut enable_w = [0.0f32; BLOCK_FRAMES];
let mut enable = self.enable;
for w in &mut enable_w {
enable += (target - enable).clamp(-step, step);
*w = enable;
}
self.enable = enable;
for ch in 0..block.channels() {
let out = block.channel_mut(ch);
for (i, sample) in out.iter_mut().enumerate() {
let weight = correction * enable_w[i];
let high = weight * self.high[ch][i] + (1.0 - weight) * self.high_raw[ch][i];
*sample = self.low[ch][i] + high;
}
}
}
fn latency_frames(&self) -> usize {
debug_assert_eq!(self.low_delay.latency_frames(), KEYLOCK_LATENCY_FRAMES);
KEYLOCK_LATENCY_FRAMES
}
fn reset(&mut self) {
self.split.reset();
self.low_delay.reset();
self.raw_high_delay.reset();
self.sola.reset();
self.enable = f32::NAN;
}
}
#[cfg(test)]
mod tests {
use super::*;
const SR: u32 = 44_100;
fn run_blocks(stage: &mut KeylockStage, input: &[f32], rate: f64) -> Vec<f32> {
run_blocks_keylock(stage, input, rate, 1.0)
}
fn run_blocks_keylock(
stage: &mut KeylockStage,
input: &[f32],
rate: f64,
keylock: f64,
) -> Vec<f32> {
let mut block = BlockBuf::new(1);
let ctx = StageCtx {
embedded_rate: rate,
embedded_rate_slope: 0.0,
onsets: &[],
modulation_hold: false,
has_artifact: false,
keylock,
};
let mut out = Vec::with_capacity(input.len());
for chunk in input.chunks_exact(BLOCK_FRAMES) {
block.channel_mut(0).copy_from_slice(chunk);
stage.process(&mut block, &ctx);
out.extend_from_slice(block.channel(0));
}
out
}
fn sine(freq: f64, len: usize, amp: f32) -> Vec<f32> {
(0..len)
.map(|i| amp * (2.0 * std::f64::consts::PI * freq * i as f64 / SR as f64).sin() as f32)
.collect()
}
#[test]
fn seam_level_recovers_after_a_nudge() {
let mut stage = KeylockStage::new(SR, 1);
let seam_hz = 170.0;
let mut phase_seam = 0.0f64;
let mut phase_hi = 0.0f64;
let mut block = BlockBuf::new(1);
let mut collected = Vec::new();
let total_secs = 8.0;
let total_blocks = (total_secs * SR as f64 / BLOCK_FRAMES as f64) as usize;
for bi in 0..total_blocks {
let t = (bi * BLOCK_FRAMES) as f64 / SR as f64;
let rate = if t < 2.0 {
1.0
} else if t < 2.2 {
1.0 + 0.04 * (t - 2.0) / 0.2
} else if t < 2.3 {
1.04
} else if t < 2.5 {
1.04 - 0.04 * (t - 2.3) / 0.2
} else {
1.0
};
for s in block.channel_mut(0).iter_mut() {
phase_seam += 2.0 * std::f64::consts::PI * seam_hz * rate / SR as f64;
phase_hi += 2.0 * std::f64::consts::PI * 880.0 * rate / SR as f64;
*s = 0.15 * phase_seam.sin() as f32 + 0.5 * phase_hi.sin() as f32;
}
let ctx = StageCtx {
embedded_rate: rate,
embedded_rate_slope: 0.0,
onsets: &[],
modulation_hold: false,
has_artifact: false,
keylock: 1.0,
};
stage.process(&mut block, &ctx);
collected.extend_from_slice(block.channel(0));
}
let goertzel = |lo: usize, hi: usize| -> f64 {
let w = 2.0 * std::f64::consts::PI * seam_hz / SR as f64;
let coeff = 2.0 * w.cos();
let (mut s1, mut s2) = (0.0f64, 0.0f64);
for &x in &collected[lo..hi] {
let s0 = x as f64 + coeff * s1 - s2;
s2 = s1;
s1 = s0;
}
(s1 * s1 + s2 * s2 - coeff * s1 * s2) / ((hi - lo) as f64 / 2.0).powi(2)
};
let sr = SR as usize;
let baseline = goertzel(sr, 2 * sr); let after = goertzel(6 * sr, 8 * sr); let loss_db = 10.0 * (after / baseline).log10();
println!(
"seam 170 Hz power: baseline {baseline:.5}, after nudge {after:.5} ({loss_db:+.2} dB)"
);
assert!(
loss_db > -1.5,
"seam level did not recover after the nudge: {loss_db:+.2} dB \
(parked SOLA drift de-phasing the bands)"
);
}
#[test]
fn keylock_holds_pitch_across_the_corrected_range() {
for (rate, label) in [(1.03f64, "dj"), (1.15, "wide")] {
let mut stage = KeylockStage::new(SR, 1);
let shifted = sine(440.0 * rate, SR as usize * 3, 0.6);
let out = run_blocks(&mut stage, &shifted, rate);
let freq = measure_freq(&out[SR as usize..SR as usize * 2]);
let cents = 1_200.0 * (freq / 440.0).log2();
assert!(
cents.abs() < 12.0,
"{label} rate: pitch off by {cents:.1} cents ({freq:.2} Hz)"
);
}
}
fn measure_freq(scan: &[f32]) -> f64 {
let (mut first, mut last, mut count) = (None, None, 0usize);
for i in 1..scan.len() {
let (a, b) = (scan[i - 1] as f64, scan[i] as f64);
if a <= 0.0 && b > 0.0 {
let t = (i - 1) as f64 + a / (a - b);
if first.is_none() {
first = Some(t);
}
last = Some(t);
count += 1;
}
}
(count - 1) as f64 * SR as f64 / (last.unwrap() - first.unwrap())
}
#[test]
fn keylock_disabled_is_delay_matched_varispeed() {
let rate = 1.06f64;
let mut stage = KeylockStage::new(SR, 1);
let shifted = sine(440.0 * rate, SR as usize * 3, 0.6);
let out = run_blocks_keylock(&mut stage, &shifted, rate, 0.0);
let scan = &out[SR as usize..SR as usize * 2];
let freq = measure_freq(scan);
let cents = 1200.0 * (freq / (440.0 * rate)).log2();
assert!(
cents.abs() < 3.0,
"bypassed output re-pitched: off by {cents:.1} cents ({freq:.2} Hz)"
);
let rms = |xs: &[f32]| {
(xs.iter().map(|&x| x as f64 * x as f64).sum::<f64>() / xs.len() as f64).sqrt()
};
let level_db = 20.0 * (rms(scan) / rms(&shifted[SR as usize..SR as usize * 2])).log10();
assert!(
level_db.abs() < 0.5,
"bypassed output level off by {level_db:+.2} dB"
);
}
#[test]
fn keylock_toggle_is_click_free_and_converges() {
let rate = 1.06f64;
let secs = SR as usize;
let shifted = sine(440.0 * rate, secs * 6, 0.6);
let mut stage = KeylockStage::new(SR, 1);
let mut out = Vec::with_capacity(shifted.len());
let mut block = BlockBuf::new(1);
for (bi, chunk) in shifted.chunks_exact(BLOCK_FRAMES).enumerate() {
let start = bi * BLOCK_FRAMES;
let keylock = if (secs * 2..secs * 4).contains(&start) {
0.0
} else {
1.0
};
let ctx = StageCtx {
embedded_rate: rate,
embedded_rate_slope: 0.0,
onsets: &[],
modulation_hold: false,
has_artifact: false,
keylock,
};
block.channel_mut(0).copy_from_slice(chunk);
stage.process(&mut block, &ctx);
out.extend_from_slice(block.channel(0));
}
let max_step = out
.windows(2)
.skip(secs / 2) .map(|w| (w[1] - w[0]).abs())
.fold(0.0f32, f32::max);
let signal_slew = 0.6 * (2.0 * std::f64::consts::PI * 440.0 * rate / SR as f64) as f32;
assert!(
max_step < signal_slew * 1.5,
"toggle clicked: max step {max_step:.4} vs signal slew {signal_slew:.4}"
);
let corrected = measure_freq(&out[secs..secs * 2]);
let bypassed = measure_freq(&out[secs * 3..secs * 4]);
let recorrected = measure_freq(&out[secs * 5..]);
let cents = |f: f64, target: f64| 1200.0 * (f / target).log2();
assert!(
cents(corrected, 440.0).abs() < 12.0,
"keylock phase off: {corrected:.2} Hz"
);
assert!(
cents(bypassed, 440.0 * rate).abs() < 12.0,
"bypass phase not at varispeed pitch: {bypassed:.2} Hz"
);
assert!(
cents(recorrected, 440.0).abs() < 12.0,
"re-enabled phase off: {recorrected:.2} Hz"
);
}
}