Skip to main content

max / audiofiles

15.5 KB · 425 lines History Blame Raw
1 //! Audio analysis pipeline: orchestrates decoding, feature extraction, classification, and DB persistence.
2 //!
3 //! The pipeline decodes audio to mono f32 via Symphonia, then runs configurable
4 //! stages (loudness, spectral features, BPM/key detection, loop detection,
5 //! classification) and persists results to the `audio_analysis` table.
6
7 pub mod basic;
8 pub mod bpm;
9 pub mod classify;
10 pub mod config;
11 pub mod decode;
12 pub mod loop_detect;
13 #[cfg(feature = "analysis")]
14 pub mod loudness;
15 pub mod mfcc;
16 pub mod spectral;
17 pub mod suggest;
18 pub mod waveform;
19 pub mod worker;
20
21 use std::path::Path;
22
23 use classify::SampleClass;
24 use config::AnalysisConfig;
25
26 use crate::db::Database;
27 use crate::error::{unix_now, CoreError};
28 use crate::fingerprint;
29 use tracing::instrument;
30
31 /// Complete analysis result for a single sample.
32 #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
33 pub struct AnalysisResult {
34 /// Content-addressed hash identifying this sample in the store.
35 pub hash: String,
36 /// Total duration in seconds.
37 pub duration: f64,
38 /// Sample rate in Hz.
39 pub sample_rate: u32,
40 /// Number of channels in the source file.
41 pub channels: u16,
42 /// Peak amplitude in dBFS.
43 pub peak_db: Option<f64>,
44 /// RMS loudness in dBFS.
45 pub rms_db: Option<f64>,
46 /// Integrated loudness in LUFS.
47 pub lufs: Option<f64>,
48 /// Estimated tempo in beats per minute.
49 pub bpm: Option<f64>,
50 /// Estimated musical key (e.g. "A minor").
51 pub musical_key: Option<String>,
52 /// Whether the sample is detected as a seamless loop.
53 pub is_loop: Option<bool>,
54 /// Spectral centroid in Hz (brightness measure).
55 pub spectral_centroid: Option<f64>,
56 /// Spectral flatness (0 = tonal, 1 = noise-like).
57 pub spectral_flatness: Option<f64>,
58 /// Spectral rolloff frequency in Hz.
59 pub spectral_rolloff: Option<f64>,
60 /// Zero-crossing rate (proportion of sign changes per sample).
61 pub zero_crossing_rate: Option<f64>,
62 /// Onset detection strength.
63 pub onset_strength: Option<f64>,
64 /// Heuristic sample classification (kick, snare, pad, etc.).
65 pub classification: Option<SampleClass>,
66 /// Peak envelope fingerprint for near-duplicate detection.
67 pub fingerprint: Option<Vec<u8>>,
68 /// Spectral bandwidth in Hz (spread of energy around centroid).
69 pub spectral_bandwidth: Option<f64>,
70 /// Variance of per-frame spectral centroids (spectral evolution).
71 pub centroid_variance: Option<f64>,
72 /// Peak-to-RMS ratio in linear domain (transient sharpness).
73 pub crest_factor: Option<f64>,
74 /// Time to 90% of peak amplitude in seconds (onset speed).
75 pub attack_time: Option<f64>,
76 /// ML classifier confidence (0.0-1.0). 0.0 when using rule-based fallback.
77 pub classification_confidence: Option<f64>,
78 }
79
80 /// Run all configured analyses on a single sample file.
81 #[instrument(skip_all)]
82 pub fn analyze_sample(
83 hash: &str,
84 path: &Path,
85 config: &AnalysisConfig,
86 ) -> Result<AnalysisResult, CoreError> {
87 // Guard against memory exhaustion: reject files over 2 GB before decoding.
88 // A 2 GB compressed file would expand to several GB of f32 samples.
89 const MAX_FILE_SIZE: u64 = 2 * 1024 * 1024 * 1024;
90 if let Ok(metadata) = std::fs::metadata(path) {
91 if metadata.len() > MAX_FILE_SIZE {
92 return Err(CoreError::Analysis(crate::error::AnalysisError::ProbeFailed(
93 format!("file too large for analysis ({} MB, max {} MB)",
94 metadata.len() / (1024 * 1024), MAX_FILE_SIZE / (1024 * 1024)),
95 )));
96 }
97 }
98
99 let decoded = decode::decode_to_mono(path)?;
100
101 // Hard cap: reject files over 30 minutes to prevent memory exhaustion.
102 // A 30-minute 96kHz mono signal is ~660 MB of f32 — beyond that is almost
103 // certainly not a sample.
104 const MAX_DECODE_DURATION: f64 = 1800.0;
105 if decoded.duration > MAX_DECODE_DURATION {
106 return Err(CoreError::Analysis(crate::error::AnalysisError::ProbeFailed(
107 format!("file too long for analysis ({:.0}s, max {MAX_DECODE_DURATION}s)", decoded.duration),
108 )));
109 }
110
111 // Cap samples for expensive analyses (STFT, BPM/key). Cheap analyses and
112 // fingerprint use the full signal.
113 let capped_samples: &[f32] = if let Some(max_secs) = config.max_analysis_seconds {
114 let max_samples = (max_secs * decoded.sample_rate as f64) as usize;
115 &decoded.samples[..decoded.samples.len().min(max_samples)]
116 } else {
117 &decoded.samples
118 };
119
120 let mut result = AnalysisResult {
121 hash: hash.to_string(),
122 duration: decoded.duration,
123 sample_rate: decoded.sample_rate,
124 channels: decoded.channels,
125 peak_db: None,
126 rms_db: None,
127 lufs: None,
128 bpm: None,
129 musical_key: None,
130 is_loop: None,
131 spectral_centroid: None,
132 spectral_flatness: None,
133 spectral_rolloff: None,
134 zero_crossing_rate: None,
135 onset_strength: None,
136 classification: None,
137 fingerprint: None,
138 spectral_bandwidth: None,
139 centroid_variance: None,
140 crest_factor: None,
141 attack_time: None,
142 classification_confidence: None,
143 };
144
145 // Basic loudness (always fast — uses full signal)
146 if config.loudness {
147 result.peak_db = Some(basic::peak_db(&decoded.samples));
148 result.rms_db = Some(basic::rms_db(&decoded.samples));
149 result.crest_factor = Some(basic::crest_factor(&decoded.samples));
150 result.attack_time = Some(basic::attack_time(&decoded.samples, decoded.sample_rate));
151 #[cfg(feature = "analysis")]
152 {
153 result.lufs =
154 Some(loudness::measure_lufs(&decoded.samples, decoded.sample_rate));
155 }
156 }
157
158 // Spectral features (uses capped samples)
159 #[cfg(feature = "analysis")]
160 if config.spectral {
161 let (features, magnitude_frames) =
162 spectral::compute_spectral_features_with_frames(capped_samples, decoded.sample_rate);
163 result.spectral_centroid = Some(features.centroid);
164 result.spectral_flatness = Some(features.flatness);
165 result.spectral_rolloff = Some(features.rolloff);
166 result.zero_crossing_rate = Some(features.zero_crossing_rate);
167 result.onset_strength = Some(features.onset_strength);
168 result.spectral_bandwidth = Some(features.bandwidth);
169 result.centroid_variance = Some(features.centroid_variance);
170
171 // Classification requires spectral + waveform features
172 if config.classify {
173 // Compute MFCCs from STFT magnitude frames
174 let mfcc_features =
175 mfcc::compute_mfccs(&magnitude_frames, decoded.sample_rate, 1024);
176
177 let input = classify::ClassifyInput::with_mfccs(
178 &features,
179 decoded.duration,
180 result.crest_factor.unwrap_or(0.0),
181 result.attack_time.unwrap_or(0.0),
182 &mfcc_features,
183 );
184
185 let ml_result = classify::classify_ml(&input);
186 result.classification = Some(ml_result.class);
187 result.classification_confidence = Some(ml_result.confidence);
188 }
189 }
190
191 // Smart skip: use classification to decide if BPM/key/loop make sense.
192 // Drums, impacts, noise, foley, ambience, and textures skip these expensive stages.
193 let (want_bpm, want_key, want_loop) = if config.smart_skip {
194 if let Some(ref class) = result.classification {
195 (
196 config.bpm && class.has_rhythm(),
197 config.key && class.has_pitch(),
198 config.loop_detect && class.has_rhythm(),
199 )
200 } else {
201 // No classification result — run everything requested
202 (config.bpm, config.key, config.loop_detect)
203 }
204 } else {
205 (config.bpm, config.key, config.loop_detect)
206 };
207
208 // BPM + key detection (uses capped samples)
209 if want_bpm || want_key {
210 let bpm_key = bpm::detect_bpm_key(capped_samples, decoded.sample_rate, 2.0);
211 if want_bpm {
212 result.bpm = bpm_key.bpm;
213 }
214 if want_key {
215 result.musical_key = bpm_key.key;
216 }
217 }
218
219 // Loop detection
220 if want_loop {
221 result.is_loop = Some(loop_detect::is_loop(
222 &decoded.samples,
223 decoded.sample_rate,
224 result.bpm,
225 ));
226 }
227
228 // Fingerprint for near-duplicate detection (uses full signal)
229 if config.fingerprint {
230 result.fingerprint = Some(fingerprint::compute_envelope(
231 &decoded.samples,
232 decoded.sample_rate,
233 ));
234 }
235
236 Ok(result)
237 }
238
239 /// Save analysis results to the database, overwriting any previous results for this hash.
240 #[instrument(skip_all)]
241 pub fn save_analysis(db: &Database, result: &AnalysisResult) -> Result<(), CoreError> {
242 let now = unix_now();
243 db.conn().execute(
244 "INSERT OR REPLACE INTO audio_analysis (
245 hash, duration, sample_rate, channels,
246 peak_db, rms_db, lufs,
247 bpm, musical_key,
248 is_loop,
249 spectral_centroid, spectral_flatness, spectral_rolloff,
250 zero_crossing_rate, onset_strength,
251 classification,
252 spectral_bandwidth, centroid_variance, crest_factor, attack_time,
253 classification_confidence,
254 analyzed_at
255 ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20, ?21, ?22)",
256 rusqlite::params![
257 result.hash,
258 result.duration,
259 result.sample_rate,
260 result.channels,
261 result.peak_db,
262 result.rms_db,
263 result.lufs,
264 result.bpm,
265 result.musical_key,
266 result.is_loop,
267 result.spectral_centroid,
268 result.spectral_flatness,
269 result.spectral_rolloff,
270 result.zero_crossing_rate,
271 result.onset_strength,
272 result.classification.as_ref().map(|c| c.as_str()),
273 result.spectral_bandwidth,
274 result.centroid_variance,
275 result.crest_factor,
276 result.attack_time,
277 result.classification_confidence,
278 now,
279 ],
280 )?;
281
282 if let Some(ref envelope) = result.fingerprint {
283 fingerprint::save_fingerprint(
284 db,
285 &fingerprint::Fingerprint {
286 hash: result.hash.clone(),
287 envelope: envelope.clone(),
288 sample_rate: result.sample_rate,
289 },
290 )?;
291 }
292
293 Ok(())
294 }
295
296 /// Load analysis results for a sample by hash. Returns `None` if no analysis exists.
297 #[instrument(skip_all)]
298 pub fn load_analysis(db: &Database, hash: &str) -> Option<AnalysisResult> {
299 db.conn()
300 .query_row(
301 "SELECT hash, duration, sample_rate, channels, peak_db, rms_db, lufs,
302 bpm, musical_key, is_loop, spectral_centroid, spectral_flatness,
303 spectral_rolloff, zero_crossing_rate, onset_strength, classification,
304 spectral_bandwidth, centroid_variance, crest_factor, attack_time,
305 classification_confidence
306 FROM audio_analysis WHERE hash = ?1",
307 [hash],
308 |row| {
309 let class_str: Option<String> = row.get(15)?;
310 Ok(AnalysisResult {
311 hash: row.get(0)?,
312 duration: row.get(1)?,
313 sample_rate: row.get(2)?,
314 channels: row.get(3)?,
315 peak_db: row.get(4)?,
316 rms_db: row.get(5)?,
317 lufs: row.get(6)?,
318 bpm: row.get(7)?,
319 musical_key: row.get(8)?,
320 is_loop: row.get(9)?,
321 spectral_centroid: row.get(10)?,
322 spectral_flatness: row.get(11)?,
323 spectral_rolloff: row.get(12)?,
324 zero_crossing_rate: row.get(13)?,
325 onset_strength: row.get(14)?,
326 classification: class_str
327 .and_then(|s| s.parse::<classify::SampleClass>().ok()),
328 spectral_bandwidth: row.get(16)?,
329 centroid_variance: row.get(17)?,
330 crest_factor: row.get(18)?,
331 attack_time: row.get(19)?,
332 classification_confidence: row.get(20)?,
333 fingerprint: None,
334 })
335 },
336 )
337 .ok()
338 }
339
340 #[cfg(test)]
341 mod tests {
342 use super::*;
343
344 #[test]
345 fn analysis_result_construction_with_defaults() {
346 let result = AnalysisResult {
347 hash: "abc123".to_string(),
348 duration: 1.5,
349 sample_rate: 48000,
350 channels: 1,
351 peak_db: None,
352 rms_db: None,
353 lufs: None,
354 bpm: None,
355 musical_key: None,
356 is_loop: None,
357 spectral_centroid: None,
358 spectral_flatness: None,
359 spectral_rolloff: None,
360 zero_crossing_rate: None,
361 onset_strength: None,
362 classification: None,
363 fingerprint: None,
364 spectral_bandwidth: None,
365 centroid_variance: None,
366 crest_factor: None,
367 attack_time: None,
368 classification_confidence: None,
369 };
370 assert_eq!(result.hash, "abc123");
371 assert_eq!(result.sample_rate, 48000);
372 assert_eq!(result.channels, 1);
373 assert!((result.duration - 1.5).abs() < f64::EPSILON);
374 assert!(result.peak_db.is_none());
375 assert!(result.bpm.is_none());
376 assert!(result.classification.is_none());
377 assert!(result.spectral_bandwidth.is_none());
378 assert!(result.crest_factor.is_none());
379 }
380
381 #[test]
382 fn analysis_result_fully_populated() {
383 let result = AnalysisResult {
384 hash: "def456".to_string(),
385 duration: 3.2,
386 sample_rate: 44100,
387 channels: 2,
388 peak_db: Some(-0.5),
389 rms_db: Some(-12.0),
390 lufs: Some(-14.0),
391 bpm: Some(128.0),
392 musical_key: Some("A minor".to_string()),
393 is_loop: Some(true),
394 spectral_centroid: Some(1500.0),
395 spectral_flatness: Some(0.15),
396 spectral_rolloff: Some(8000.0),
397 zero_crossing_rate: Some(0.05),
398 onset_strength: Some(30.0),
399 classification: Some(SampleClass::Kick),
400 fingerprint: None,
401 spectral_bandwidth: Some(2500.0),
402 centroid_variance: Some(50000.0),
403 crest_factor: Some(4.5),
404 attack_time: Some(0.005),
405 classification_confidence: Some(0.87),
406 };
407 assert_eq!(result.peak_db, Some(-0.5));
408 assert_eq!(result.rms_db, Some(-12.0));
409 assert_eq!(result.lufs, Some(-14.0));
410 assert_eq!(result.bpm, Some(128.0));
411 assert_eq!(result.musical_key.as_deref(), Some("A minor"));
412 assert_eq!(result.is_loop, Some(true));
413 assert_eq!(result.spectral_centroid, Some(1500.0));
414 assert_eq!(result.spectral_flatness, Some(0.15));
415 assert_eq!(result.spectral_rolloff, Some(8000.0));
416 assert_eq!(result.zero_crossing_rate, Some(0.05));
417 assert_eq!(result.onset_strength, Some(30.0));
418 assert_eq!(result.classification, Some(SampleClass::Kick));
419 assert_eq!(result.spectral_bandwidth, Some(2500.0));
420 assert_eq!(result.centroid_variance, Some(50000.0));
421 assert_eq!(result.crest_factor, Some(4.5));
422 assert_eq!(result.attack_time, Some(0.005));
423 }
424 }
425