/* * SPDX-FileCopyrightText: 2020 Stalwart Labs LLC * * SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL */ use rand::{RngExt, seq::IndexedRandom}; #[derive(rkyv::Archive, rkyv::Deserialize, rkyv::Serialize, Debug)] pub struct SampleReservoir { pub spam: SampleReservoirClass, pub ham: SampleReservoirClass, } #[derive(rkyv::Archive, rkyv::Deserialize, rkyv::Serialize, Debug)] pub struct SampleReservoirClass { pub buffer: Vec, pub total_seen: u64, } impl SampleReservoir { pub fn update_reservoir(&mut self, item: &T, is_spam: bool, capacity: usize) { let class = if is_spam { &mut self.spam } else { &mut self.ham }; class.total_seen += 1; if class.buffer.len() < capacity { class.buffer.push(item.clone()); } else if let Some(buf) = class .buffer .get_mut(rand::rng().random_range(0..class.total_seen as usize)) { *buf = item.clone(); } } pub fn update_counts(&mut self, is_spam: bool) { let class = if is_spam { &mut self.spam } else { &mut self.ham }; class.total_seen += 1; } pub fn replay_samples( &mut self, count_needed: usize, is_spam: bool, ) -> impl Iterator { (if is_spam { &mut self.spam } else { &mut self.ham }) .buffer .sample(&mut rand::rng(), count_needed) } pub fn remove_sample(&mut self, item: &T, is_spam: bool) { let class = if is_spam { &mut self.spam } else { &mut self.ham }; if let Some(pos) = class.buffer.iter().position(|x| x == item) { class.buffer.swap_remove(pos); } } } impl Default for SampleReservoir { fn default() -> Self { SampleReservoir { spam: SampleReservoirClass { buffer: Vec::new(), total_seen: 0, }, ham: SampleReservoirClass { buffer: Vec::new(), total_seen: 0, }, } } }