use super::hash::Hash;
use super::matching::FingerprintMatcher;
use super::Fingerprint;
use std::collections::HashMap;
pub struct FingerprintDatabase {
hash_index: HashMap<Hash, Vec<(String, f64)>>,
tracks: HashMap<String, TrackMetadata>,
matcher: FingerprintMatcher,
}
impl FingerprintDatabase {
#[must_use]
pub fn new() -> Self {
Self {
hash_index: HashMap::new(),
tracks: HashMap::new(),
matcher: FingerprintMatcher::default(),
}
}
#[must_use]
pub fn with_matcher(matcher: FingerprintMatcher) -> Self {
Self {
hash_index: HashMap::new(),
tracks: HashMap::new(),
matcher,
}
}
pub fn add_fingerprint(&mut self, track_id: impl Into<String>, fingerprint: Fingerprint) {
let track_id = track_id.into();
self.tracks.insert(
track_id.clone(),
TrackMetadata {
id: track_id.clone(),
duration: fingerprint.duration,
sample_rate: fingerprint.sample_rate,
hash_count: fingerprint.hashes.len(),
},
);
for (hash, time) in fingerprint.hashes {
self.hash_index
.entry(hash)
.or_default()
.push((track_id.clone(), time));
}
}
pub fn remove_fingerprint(&mut self, track_id: &str) -> bool {
if self.tracks.remove(track_id).is_none() {
return false;
}
for entries in self.hash_index.values_mut() {
entries.retain(|(id, _)| id != track_id);
}
self.hash_index.retain(|_, entries| !entries.is_empty());
true
}
#[must_use]
pub fn find_matches(&self, query: &Fingerprint, min_confidence: f64) -> Vec<Match> {
let mut candidate_scores: HashMap<String, Vec<(Hash, f64, f64)>> = HashMap::new();
for (query_hash, query_time) in &query.hashes {
if let Some(entries) = self.hash_index.get(query_hash) {
for (track_id, ref_time) in entries {
candidate_scores.entry(track_id.clone()).or_default().push((
*query_hash,
*query_time,
*ref_time,
));
}
}
}
let mut matches = Vec::new();
for (track_id, hash_matches) in candidate_scores {
if let Some(metadata) = self.tracks.get(&track_id) {
let ref_hashes: Vec<(Hash, f64)> = hash_matches
.iter()
.map(|(h, _, ref_time)| (*h, *ref_time))
.collect();
let ref_fingerprint = Fingerprint::new(
ref_hashes,
metadata.sample_rate,
metadata.duration,
query.config.clone(),
);
if let Some(result) = self.matcher.match_fingerprint(query, &ref_fingerprint) {
if result.confidence >= min_confidence && self.matcher.verify_match(&result) {
matches.push(Match {
track_id: track_id.clone(),
confidence: result.confidence,
time_offset: result.time_offset,
match_count: result.match_count,
query_coverage: result.query_coverage(),
reference_coverage: result.reference_coverage(),
});
}
}
}
}
matches.sort_by(|a, b| {
b.confidence
.partial_cmp(&a.confidence)
.unwrap_or(std::cmp::Ordering::Equal)
});
matches
}
#[must_use]
pub fn find_best_match(&self, query: &Fingerprint, min_confidence: f64) -> Option<Match> {
self.find_matches(query, min_confidence).into_iter().next()
}
#[must_use]
pub fn contains_track(&self, track_id: &str) -> bool {
self.tracks.contains_key(track_id)
}
#[must_use]
pub fn get_track_metadata(&self, track_id: &str) -> Option<&TrackMetadata> {
self.tracks.get(track_id)
}
#[must_use]
pub fn track_ids(&self) -> Vec<String> {
self.tracks.keys().cloned().collect()
}
#[must_use]
pub fn track_count(&self) -> usize {
self.tracks.len()
}
#[must_use]
pub fn hash_count(&self) -> usize {
self.hash_index.len()
}
#[must_use]
pub fn statistics(&self) -> DatabaseStatistics {
let total_entries: usize = self.hash_index.values().map(Vec::len).sum();
let avg_collisions = if !self.hash_index.is_empty() {
total_entries as f64 / self.hash_index.len() as f64
} else {
0.0
};
let total_duration: f64 = self.tracks.values().map(|m| m.duration).sum();
DatabaseStatistics {
track_count: self.tracks.len(),
unique_hash_count: self.hash_index.len(),
total_hash_entries: total_entries,
avg_hash_collisions: avg_collisions,
total_duration,
}
}
pub fn clear(&mut self) {
self.hash_index.clear();
self.tracks.clear();
}
pub fn merge(&mut self, other: Self) {
self.tracks.extend(other.tracks);
for (hash, entries) in other.hash_index {
self.hash_index.entry(hash).or_default().extend(entries);
}
}
pub fn optimize(&mut self) {
for entries in self.hash_index.values_mut() {
entries.sort_by(|a, b| {
a.0.cmp(&b.0)
.then(a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
});
entries.dedup();
}
self.hash_index.retain(|_, entries| !entries.is_empty());
let mut hash_counts: HashMap<String, usize> = HashMap::new();
for entries in self.hash_index.values() {
for (track_id, _) in entries {
*hash_counts.entry(track_id.clone()).or_insert(0) += 1;
}
}
for (track_id, metadata) in &mut self.tracks {
if let Some(&count) = hash_counts.get(track_id) {
metadata.hash_count = count;
}
}
}
#[must_use]
pub fn export(&self) -> DatabaseExport {
let mut entries = Vec::new();
for (hash, track_times) in &self.hash_index {
for (track_id, time) in track_times {
entries.push(HashEntry {
hash: hash.value(),
track_id: track_id.clone(),
time: *time,
});
}
}
DatabaseExport {
tracks: self.tracks.values().cloned().collect(),
entries,
}
}
pub fn import(export: DatabaseExport) -> Self {
let mut db = Self::new();
for metadata in export.tracks {
db.tracks.insert(metadata.id.clone(), metadata);
}
for entry in export.entries {
let hash = Hash::from(entry.hash);
db.hash_index
.entry(hash)
.or_default()
.push((entry.track_id, entry.time));
}
db
}
#[must_use]
pub fn find_duplicates(&self, min_confidence: f64) -> Vec<(String, String, f64)> {
let mut duplicates = Vec::new();
let track_ids: Vec<_> = self.tracks.keys().collect();
for i in 0..track_ids.len() {
for j in (i + 1)..track_ids.len() {
let id1 = track_ids[i];
let id2 = track_ids[j];
let hashes1 = self.get_track_hashes(id1);
let hashes2 = self.get_track_hashes(id2);
if let (Some(meta1), Some(meta2)) = (self.tracks.get(id1), self.tracks.get(id2)) {
let fp1 = Fingerprint::new(
hashes1,
meta1.sample_rate,
meta1.duration,
Default::default(),
);
let fp2 = Fingerprint::new(
hashes2,
meta2.sample_rate,
meta2.duration,
Default::default(),
);
if let Some(result) = self.matcher.match_fingerprint(&fp1, &fp2) {
if result.confidence >= min_confidence {
duplicates.push((id1.clone(), id2.clone(), result.confidence));
}
}
}
}
}
duplicates
}
fn get_track_hashes(&self, track_id: &str) -> Vec<(Hash, f64)> {
let mut hashes = Vec::new();
for (hash, entries) in &self.hash_index {
for (id, time) in entries {
if id == track_id {
hashes.push((*hash, *time));
}
}
}
hashes.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
hashes
}
}
impl Default for FingerprintDatabase {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone, Debug)]
pub struct TrackMetadata {
pub id: String,
pub duration: f64,
pub sample_rate: u32,
pub hash_count: usize,
}
#[derive(Clone, Debug)]
pub struct Match {
pub track_id: String,
pub confidence: f64,
pub time_offset: f64,
pub match_count: usize,
pub query_coverage: f64,
pub reference_coverage: f64,
}
impl Match {
#[must_use]
pub fn is_strong_match(&self) -> bool {
self.confidence >= 0.7 && self.query_coverage >= 0.5
}
#[must_use]
pub fn is_likely_duplicate(&self) -> bool {
self.confidence >= 0.9 && self.query_coverage >= 0.8
}
#[must_use]
pub fn is_partial_match(&self) -> bool {
self.confidence >= 0.3 && self.query_coverage < 0.5
}
}
#[derive(Clone, Debug, Default)]
pub struct DatabaseStatistics {
pub track_count: usize,
pub unique_hash_count: usize,
pub total_hash_entries: usize,
pub avg_hash_collisions: f64,
pub total_duration: f64,
}
impl DatabaseStatistics {
#[must_use]
pub fn estimated_size_bytes(&self) -> usize {
self.total_hash_entries * (8 + 32 + 8)
}
#[must_use]
pub fn avg_hashes_per_track(&self) -> f64 {
if self.track_count > 0 {
self.total_hash_entries as f64 / self.track_count as f64
} else {
0.0
}
}
}
#[derive(Clone, Debug)]
pub struct DatabaseExport {
pub tracks: Vec<TrackMetadata>,
pub entries: Vec<HashEntry>,
}
#[derive(Clone, Debug)]
pub struct HashEntry {
pub hash: u64,
pub track_id: String,
pub time: f64,
}
pub struct FingerprintCache {
cache: HashMap<String, Fingerprint>,
max_size: usize,
}
impl FingerprintCache {
#[must_use]
pub fn new(max_size: usize) -> Self {
Self {
cache: HashMap::new(),
max_size,
}
}
#[must_use]
pub fn get(&self, track_id: &str) -> Option<&Fingerprint> {
self.cache.get(track_id)
}
pub fn insert(&mut self, track_id: String, fingerprint: Fingerprint) {
if self.cache.len() >= self.max_size {
if let Some(key) = self.cache.keys().next().cloned() {
self.cache.remove(&key);
}
}
self.cache.insert(track_id, fingerprint);
}
pub fn clear(&mut self) {
self.cache.clear();
}
#[must_use]
pub fn size(&self) -> usize {
self.cache.len()
}
}