use alloc::vec::Vec;
use crate::{AudioBuffer, Result, TimestampMs};
#[must_use]
pub trait Fingerprinter {
type Output;
type Config: Clone + Send + Sync;
fn name(&self) -> &'static str;
fn config(&self) -> &Self::Config;
fn required_sample_rate(&self) -> u32;
fn min_samples(&self) -> usize;
fn extract(&mut self, audio: AudioBuffer<'_>) -> Result<Self::Output>;
}
#[must_use]
pub trait StreamingFingerprinter {
type Frame;
fn required_sample_rate(&self) -> u32 {
0
}
fn push(&mut self, samples: &[f32]) -> Vec<(TimestampMs, Self::Frame)>;
fn flush(&mut self) -> Vec<(TimestampMs, Self::Frame)>;
fn latency_ms(&self) -> u32;
fn push_with<F>(&mut self, samples: &[f32], mut callback: F) -> usize
where
F: FnMut(TimestampMs, &Self::Frame),
{
let frames = self.push(samples);
let n = frames.len();
for (t, frame) in frames {
callback(t, &frame);
}
n
}
fn flush_with<F>(&mut self, mut callback: F) -> usize
where
F: FnMut(TimestampMs, &Self::Frame),
{
let frames = self.flush();
let n = frames.len();
for (t, frame) in frames {
callback(t, &frame);
}
n
}
}
#[cfg(feature = "rayon")]
#[must_use]
pub fn fingerprint_batch_parallel<F, T>(
items: Vec<(T, Vec<f32>, crate::SampleRate)>,
make_fingerprinter: impl Fn() -> F + Sync,
) -> Vec<(T, Result<F::Output>)>
where
F: Fingerprinter + Send,
F::Output: Send,
T: Send,
{
use rayon::prelude::*;
items
.into_par_iter()
.map(|(tag, samples, rate)| {
let mut fp = make_fingerprinter();
let buf = AudioBuffer::new(&samples, rate);
(tag, fp.extract(buf))
})
.collect()
}
#[cfg(test)]
mod tests {
use alloc::vec;
use super::*;
struct CountByThree {
count: u32,
buffered: Vec<u32>,
}
impl CountByThree {
fn new() -> Self {
Self {
count: 0,
buffered: Vec::new(),
}
}
}
impl StreamingFingerprinter for CountByThree {
type Frame = u32;
fn push(&mut self, samples: &[f32]) -> Vec<(TimestampMs, u32)> {
let mut out = Vec::new();
for _ in samples {
self.count += 1;
if self.count.is_multiple_of(3) {
out.push((TimestampMs(self.count as u64), self.count));
}
}
self.buffered.extend(out.iter().map(|(_, f)| *f));
out
}
fn flush(&mut self) -> Vec<(TimestampMs, u32)> {
let pending: Vec<u32> = self.buffered.drain(..).collect();
pending
.into_iter()
.map(|v| (TimestampMs(v as u64 + 100), v))
.collect()
}
fn latency_ms(&self) -> u32 {
0
}
}
#[test]
fn push_with_default_impl_matches_push() {
let samples = vec![0.0_f32; 10];
let mut fp = CountByThree::new();
let mut a: Vec<(TimestampMs, u32)> = Vec::new();
a.extend(fp.push(&samples));
a.extend(fp.push(&[]));
let mut fp = CountByThree::new();
let mut b: Vec<(TimestampMs, u32)> = Vec::new();
fp.push_with(&samples, |t, f| b.push((t, *f)));
fp.push_with(&[], |t, f| b.push((t, *f)));
assert_eq!(
a.len(),
b.len(),
"push_with must call back in the same order as push yields"
);
assert_eq!(a, b, "push_with must mirror push output exactly");
}
#[test]
fn flush_with_default_impl_matches_flush() {
let mut fp = CountByThree::new();
let samples = vec![0.0_f32; 9];
let _ = fp.push(&samples);
let pending: Vec<_> = fp.flush();
let mut fp = CountByThree::new();
let _ = fp.push(&samples);
let mut collected = Vec::new();
let n = fp.flush_with(|t, f| collected.push((t, *f)));
assert_eq!(n, pending.len());
assert_eq!(collected, pending, "flush_with must mirror flush");
}
#[cfg(feature = "rayon")]
#[test]
fn batch_parallel_produces_same_results_as_sequential() {
use super::fingerprint_batch_parallel;
use crate::SampleRate;
struct Sum;
impl Fingerprinter for Sum {
type Output = f32;
type Config = ();
fn name(&self) -> &'static str {
"sum"
}
fn config(&self) -> &Self::Config {
&()
}
fn required_sample_rate(&self) -> u32 {
8_000
}
fn min_samples(&self) -> usize {
1
}
fn extract(&mut self, audio: AudioBuffer<'_>) -> crate::Result<f32> {
Ok(audio.samples.iter().sum())
}
}
let items: Vec<(u32, Vec<f32>, SampleRate)> = (0..100)
.map(|i| (i, vec![i as f32; 10], SampleRate::HZ_8000))
.collect();
let mut sequential = Vec::new();
for (tag, samples, rate) in &items {
let mut fp = Sum;
let buf = AudioBuffer::new(samples, *rate);
sequential.push((*tag, fp.extract(buf)));
}
let parallel = fingerprint_batch_parallel(items, || Sum);
assert_eq!(sequential.len(), parallel.len());
for (s, p) in sequential.iter().zip(parallel.iter()) {
assert_eq!(s.0, p.0, "tags must match in order");
assert!((s.1.as_ref().unwrap() - p.1.as_ref().unwrap()).abs() < 1e-6);
}
}
#[test]
fn push_with_reports_emitted_count() {
let mut fp = CountByThree::new();
let n = fp.push_with(&[0.0_f32; 9], |_, _| {});
assert_eq!(n, 3);
let n = fp.push_with(&[0.0_f32; 2], |_, _| {});
assert_eq!(n, 0);
}
}