use crate::common::ResizeFactor;
use crate::error::Error;
use crate::error::ErrorKind;
use crate::hash::check_seed_hash;
use crate::thetacommon::EntrySketch;
use crate::thetacommon::SketchEntry;
use crate::thetacommon::SketchScalars;
use crate::thetacommon::constants::MAX_THETA;
use crate::thetacommon::hash_table::CompactSketchParts;
use crate::thetacommon::hash_table::SketchHashTable;
pub trait UnionMergePolicy<E> {
fn merge(&self, existing: &mut E, incoming: E);
}
#[derive(Debug)]
pub struct UnionState<E, P> {
table: SketchHashTable<E>,
policy: P,
union_theta: u64,
}
impl<E, P> UnionState<E, P>
where
E: SketchEntry,
{
pub fn new(
lg_k: u8,
resize_factor: ResizeFactor,
sampling_probability: f32,
seed: u64,
policy: P,
) -> Self {
let table = SketchHashTable::new(lg_k, resize_factor, sampling_probability, seed);
Self {
union_theta: table.theta(),
table,
policy,
}
}
pub fn update<S>(&mut self, sketch: S) -> Result<(), Error>
where
S: EntrySketch<Entry = E>,
P: UnionMergePolicy<E>,
{
let SketchScalars {
seed_hash,
theta,
empty,
ordered,
..
} = sketch.scalars();
if empty {
return Ok(());
}
check_seed_hash(
self.table.seed_hash(),
seed_hash,
"union update",
ErrorKind::InvalidArgument,
)?;
self.table.set_empty(false);
self.union_theta = self.union_theta.min(theta);
for entry in sketch.entries() {
let hash = entry.hash();
if hash < self.union_theta && hash < self.table.theta() {
self.table.upsert_entry(hash, |existing| match existing {
Some(existing) => {
self.policy.merge(existing, entry);
None
}
None => Some(entry),
});
} else if ordered {
break;
}
}
self.union_theta = self.union_theta.min(self.table.theta());
Ok(())
}
pub fn to_compact_parts(&self, ordered: bool) -> CompactSketchParts<E>
where
E: Clone,
{
let seed_hash = self.table.seed_hash();
if self.table.is_empty() {
return CompactSketchParts {
entries: vec![],
theta: self.union_theta,
seed_hash,
ordered: true,
empty: true,
};
}
let mut theta = self.union_theta.min(self.table.theta());
let mut entries = if self.union_theta >= self.table.theta() {
self.table.iter_entries().cloned().collect::<Vec<_>>()
} else {
self.table
.iter_entries()
.filter(|entry| entry.hash() < theta)
.cloned()
.collect::<Vec<_>>()
};
let nominal_num = 1usize << self.table.lg_nom_size();
if entries.len() > nominal_num {
let (_, kth, _) = entries.select_nth_unstable_by_key(nominal_num, |entry| entry.hash());
theta = kth.hash();
entries.truncate(nominal_num);
}
let ordered = ordered || (entries.len() == 1 && theta == MAX_THETA);
if ordered {
entries.sort_unstable_by_key(SketchEntry::hash);
}
CompactSketchParts {
entries,
theta,
seed_hash,
ordered,
empty: false,
}
}
pub fn reset(&mut self) {
self.table.reset();
self.union_theta = self.table.theta();
}
pub fn estimated_size(&self) -> usize {
self.table.estimated_size()
}
}