use crate::config::SsirConfig;
use crate::detection::DetectedReflection;
use crate::types::RirSegment;
use rayon::prelude::*;
pub(crate) fn build_segments(
rir: &[f32],
direct_sound_toa: usize,
direct_sound_doa: Option<[f32; 3]>,
reflections: &[DetectedReflection],
mixing_time_samples: usize,
config: &SsirConfig,
) -> Vec<RirSegment> {
if rir.is_empty() || direct_sound_toa >= rir.len() {
return Vec::new();
}
let onset_window = config.onset_window_samples();
let min_segment = config.min_segment_samples();
let final_segment = config.final_segment_samples();
let mut events: Vec<(usize, f64, Option<[f32; 3]>, bool)> = Vec::new();
let ds_energy = (rir[direct_sound_toa] as f64).powi(2);
events.push((direct_sound_toa, ds_energy, direct_sound_doa, true));
for r in reflections {
events.push((r.toa_sample, r.peak_energy, r.doa, false));
}
let mut onsets: Vec<usize> = Vec::with_capacity(events.len() + 1);
onsets.push(0);
onsets.extend(
events[1..]
.par_iter()
.map(|event| find_onset(rir, event.0, onset_window))
.collect::<Vec<_>>(),
);
let mut refined_onsets: Vec<usize> = vec![onsets[0]];
let mut refined_events: Vec<usize> = vec![0];
for (i, &onset) in onsets.iter().enumerate().skip(1) {
let prev_onset = *refined_onsets.last().unwrap();
let duration = onset.saturating_sub(prev_onset);
if i > 0 && duration < min_segment {
continue;
}
refined_onsets.push(onsets[i]);
refined_events.push(i);
}
let num_segments = refined_onsets.len();
let mut segments: Vec<RirSegment> = Vec::with_capacity(num_segments);
for seg_idx in 0..num_segments {
let event_idx = refined_events[seg_idx];
let (toa, peak_energy, doa, is_direct) = &events[event_idx];
let onset_sample = refined_onsets[seg_idx];
let end_sample = if seg_idx + 1 < num_segments {
refined_onsets[seg_idx + 1]
} else {
(onset_sample + final_segment)
.min(mixing_time_samples + final_segment)
.min(rir.len())
};
segments.push(RirSegment {
onset_sample,
end_sample,
toa_sample: *toa,
doa: *doa,
peak_energy: *peak_energy,
is_direct_sound: *is_direct,
});
}
segments
}
fn find_onset(rir: &[f32], toa: usize, onset_window: usize) -> usize {
let start = toa.saturating_sub(onset_window);
if start >= toa || toa >= rir.len() {
return toa;
}
let window: Vec<f64> = (start..=toa.min(rir.len() - 1))
.map(|i| (rir[i] as f64).powi(2))
.collect();
if window.is_empty() {
return toa;
}
let peak_energy = *window.last().unwrap();
if peak_energy < 1e-20 {
return toa;
}
let threshold = peak_energy * 0.1;
for (i, &e) in window.iter().enumerate() {
if e >= threshold {
return start + i;
}
}
toa
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_find_onset_places_before_peak() {
let mut rir = vec![0.0001f32; 200];
rir[95] = 0.05;
rir[96] = 0.1;
rir[97] = 0.3;
rir[98] = 0.6;
rir[99] = 0.9;
rir[100] = 1.0;
let onset = find_onset(&rir, 100, 24); assert!(onset <= 100, "onset should be at or before TOA");
assert!(
onset >= 95,
"onset should be within the ramp-up, got {onset}"
);
}
#[test]
fn test_build_segments_basic() {
let mut rir = vec![0.0001f32; 2400]; rir[48] = 1.0; rir[288] = 0.5; rir[480] = 0.3;
let reflections = vec![
DetectedReflection {
toa_sample: 288,
peak_energy: 0.25,
doa: None,
},
DetectedReflection {
toa_sample: 480,
peak_energy: 0.09,
doa: None,
},
];
let config = SsirConfig {
sample_rate: 48000.0,
mixing_time_ms: Some(40.0),
..SsirConfig::default()
};
let segments = build_segments(
&rir,
48,
None,
&reflections,
config.mixing_time_samples(),
&config,
);
assert_eq!(segments.len(), 3);
assert!(segments[0].is_direct_sound);
assert!(!segments[1].is_direct_sound);
assert!(!segments[2].is_direct_sound);
assert_eq!(segments[0].end_sample, segments[1].onset_sample);
assert_eq!(segments[1].end_sample, segments[2].onset_sample);
}
#[test]
fn test_short_segments_are_merged() {
let mut rir = vec![0.0001f32; 2400];
rir[48] = 1.0;
rir[288] = 0.5;
rir[290] = 0.4;
let reflections = vec![
DetectedReflection {
toa_sample: 288,
peak_energy: 0.25,
doa: None,
},
DetectedReflection {
toa_sample: 290,
peak_energy: 0.16,
doa: None,
},
];
let config = SsirConfig {
sample_rate: 48000.0,
min_segment_ms: 1.0, mixing_time_ms: Some(40.0),
..SsirConfig::default()
};
let segments = build_segments(
&rir,
48,
None,
&reflections,
config.mixing_time_samples(),
&config,
);
assert_eq!(segments.len(), 2);
}
#[test]
fn test_first_reflection_merged_when_too_short() {
let mut rir = vec![0.0001f32; 2400];
rir[100] = 1.0; rir[110] = 0.5;
let reflections = vec![DetectedReflection {
toa_sample: 110,
peak_energy: 0.25,
doa: None,
}];
let config = SsirConfig {
sample_rate: 48000.0,
direct_sound_window_ms: (0.1, 0.1), min_segment_ms: 6.0, mixing_time_ms: Some(40.0),
..SsirConfig::default()
};
let segments = build_segments(
&rir,
100,
None,
&reflections,
config.mixing_time_samples(),
&config,
);
assert_eq!(
segments.len(),
1,
"first reflection should be merged when its segment is shorter than min_segment"
);
}
}