Files
inbuxa-server/crates/nlp/src/classifier/reservoir.rs
T
jcoffey-dev 7dae9b29fd Import upstream v0.16.22, stripped
Upstream commit: 474dd0229cb20cf513036619781ed97bd8073c3f
Enterprise-only files removed or emptied: 63
Enterprise-only snippets removed: 117 in 50 files
Dangling module declarations removed: 5
Cargo edits turning enterprise off: 14
Verification: clean
Enterprise feature gates left for rebuilt features: 19 in 18 files

Produced by tools/fork/strip.py. The full report is in docs/fork/strip-reports/ on main.
2026-09-18 10:21:56 -07:00

92 lines
2.2 KiB
Rust

/*
* SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <[email protected]>
*
* 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<T> {
pub spam: SampleReservoirClass<T>,
pub ham: SampleReservoirClass<T>,
}
#[derive(rkyv::Archive, rkyv::Deserialize, rkyv::Serialize, Debug)]
pub struct SampleReservoirClass<T> {
pub buffer: Vec<T>,
pub total_seen: u64,
}
impl<T: Clone + Eq> SampleReservoir<T> {
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<Item = &T> {
(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<T> Default for SampleReservoir<T> {
fn default() -> Self {
SampleReservoir {
spam: SampleReservoirClass {
buffer: Vec::new(),
total_seen: 0,
},
ham: SampleReservoirClass {
buffer: Vec::new(),
total_seen: 0,
},
}
}
}