1use std::collections::BTreeMap;
19use std::sync::Arc;
20
21use crate::instrument::Note;
22use crate::render;
23use crate::runtime::AudioSource;
24
25const GM_DRUMS: std::ops::RangeInclusive<u8> = 35..=51;
27
28pub struct DrumKit {
30 pieces: BTreeMap<u8, Arc<Vec<f32>>>,
32 voices: Vec<DrumVoice>,
33 max_voices: usize,
35 fade_step: f32,
37}
38
39struct DrumVoice {
40 samples: Arc<Vec<f32>>,
41 pos: usize,
42 gain: f32,
43 fade: f32,
46 fade_step: f32,
47}
48
49fn render_drum(midi: u8, sample_rate: u32) -> Vec<f32> {
52 let doc = serde_json::json!({
53 "name": "drum", "duration": 1.5, "sample_rate": sample_rate, "engine": 2,
54 "root": { "type": "seq", "bpm": 120, "steps_per_beat": 4, "wave": "kit",
55 "env": { "a": 0.001, "d": 0.4, "s": 0.0, "r": 0.15 },
56 "notes": [ { "step": 0, "len": 2, "pitch": format!("midi:{midi}"), "gain": 1.0 } ] }
57 });
58 let Ok(doc) = serde_json::from_value(doc) else {
59 return Vec::new();
60 };
61 let mut s = render::render(&doc);
62 let end = s
63 .iter()
64 .rposition(|x| x.abs() > 1e-4)
65 .map(|i| i + 1)
66 .unwrap_or(0);
67 s.truncate(end);
68 s
69}
70
71impl DrumKit {
72 pub fn general_midi(sample_rate: u32) -> Self {
74 let mut pieces = BTreeMap::new();
75 for midi in GM_DRUMS {
76 let samples = render_drum(midi, sample_rate);
77 if !samples.is_empty() {
78 pieces.insert(midi, Arc::new(samples));
79 }
80 }
81 DrumKit {
82 pieces,
83 voices: Vec::new(),
84 max_voices: 32,
85 fade_step: 1.0 / (sample_rate as f32 * 0.005),
86 }
87 }
88
89 pub fn with_max_voices(mut self, max: usize) -> Self {
92 self.max_voices = max.max(1);
93 self
94 }
95
96 pub fn note_on(&mut self, note: Note, velocity: f32) {
99 if let Some(samples) = self.pieces.get(¬e.midi()) {
100 let sounding = self.voices.iter().filter(|v| v.fade_step == 0.0).count();
105 if sounding >= self.max_voices
106 && let Some(oldest) = self.voices.iter_mut().find(|v| v.fade_step == 0.0)
107 {
108 oldest.fade_step = self.fade_step;
109 }
110 if self.voices.len() >= self.max_voices * 2 {
111 self.voices.remove(0);
112 }
113 self.voices.push(DrumVoice {
114 samples: samples.clone(),
115 pos: 0,
116 gain: velocity.clamp(0.0, 1.0),
117 fade: 1.0,
118 fade_step: 0.0,
119 });
120 }
121 }
122
123 pub fn notes(&self) -> impl Iterator<Item = u8> + '_ {
125 self.pieces.keys().copied()
126 }
127
128 pub fn active_voices(&self) -> usize {
130 self.voices.len()
131 }
132}
133
134impl AudioSource for DrumKit {
135 fn fill(&mut self, out: &mut [f32]) -> usize {
136 let frames = out.len() / 2;
137 out.fill(0.0);
138 for v in self.voices.iter_mut() {
139 for f in 0..frames {
140 let Some(&s) = v.samples.get(v.pos) else {
141 break;
142 };
143 if v.fade_step > 0.0 {
144 v.fade = (v.fade - v.fade_step).max(0.0);
145 if v.fade == 0.0 {
146 v.pos = v.samples.len(); break;
148 }
149 }
150 let x = s * v.gain * v.fade;
151 out[f * 2] += x;
152 out[f * 2 + 1] += x;
153 v.pos += 1;
154 }
155 }
156 self.voices.retain(|v| v.pos < v.samples.len());
157 frames
158 }
159}
160
161#[cfg(test)]
162mod tests {
163 use super::*;
164
165 fn peak(s: &[f32]) -> f32 {
166 s.iter().fold(0.0f32, |m, &x| m.max(x.abs()))
167 }
168
169 #[test]
170 fn gm_kit_maps_and_plays_drums() {
171 let mut kit = DrumKit::general_midi(48_000);
172 assert!(kit.notes().count() >= 8, "a usable set of GM drums");
173 kit.note_on(Note(36), 1.0); kit.note_on(Note(38), 0.9); assert_eq!(kit.active_voices(), 2);
176 let mut out = vec![0.0f32; 512 * 2];
177 assert_eq!(kit.fill(&mut out), 512);
178 assert!(peak(&out) > 0.0, "the kit makes sound");
179 assert!((0..512).all(|f| out[f * 2] == out[f * 2 + 1]), "centered");
180 }
181
182 #[test]
183 fn one_shots_cull_when_finished() {
184 let mut kit = DrumKit::general_midi(48_000);
185 kit.note_on(Note(42), 1.0); assert_eq!(kit.active_voices(), 1);
187 for _ in 0..200 {
189 kit.fill(&mut vec![0.0f32; 512 * 2]);
190 }
191 assert_eq!(kit.active_voices(), 0, "finished one-shot reclaimed");
192 }
193
194 #[test]
195 fn unmapped_note_is_ignored() {
196 let mut kit = DrumKit::general_midi(48_000);
197 kit.note_on(Note(0), 1.0); assert_eq!(kit.active_voices(), 0);
199 }
200
201 #[test]
202 fn stealing_fades_the_oldest_hit_instead_of_cutting() {
203 let max_jump = |kit: &mut DrumKit, prev: f32| {
205 let mut buf = vec![0.0f32; 512 * 2];
206 kit.fill(&mut buf);
207 let mut m = 0.0f32;
208 let mut p = prev;
209 for f in 0..512 {
210 m = m.max((buf[f * 2] - p).abs());
211 p = buf[f * 2];
212 }
213 m
214 };
215 let warmup = |kit: &mut DrumKit| {
216 kit.note_on(Note(36), 1.0); let mut out = vec![0.0f32; 64 * 2];
218 kit.fill(&mut out);
219 out[out.len() - 2]
220 };
221 let mut natural = DrumKit::general_midi(48_000);
222 let prev = warmup(&mut natural);
223 let natural_jump = max_jump(&mut natural, prev);
224
225 let mut kit = DrumKit::general_midi(48_000);
228 kit.max_voices = 2;
229 let prev = warmup(&mut kit);
230 kit.note_on(Note(38), 0.0);
231 kit.note_on(Note(42), 0.0); let stolen_jump = max_jump(&mut kit, prev);
233
234 assert!(
237 stolen_jump <= natural_jump + 0.05,
238 "steal clicked: jump {stolen_jump} vs natural {natural_jump}"
239 );
240 kit.fill(&mut vec![0.0f32; 512 * 2]);
242 assert_eq!(kit.active_voices(), 2, "faded voice culled after steal");
243 }
244}