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::HASH_TABLE_REBUILD_THRESHOLD;
use crate::thetacommon::constants::MAX_THETA;
use crate::thetacommon::hash_table::CompactSketchParts;
use crate::thetacommon::hash_table::SketchHashTable;
pub trait IntersectionMergePolicy<E> {
fn merge(&self, existing: &mut E, incoming: E);
}
#[derive(Debug)]
pub struct IntersectionState<E, P> {
table: SketchHashTable<E>,
policy: P,
has_result: bool,
}
impl<E, P> IntersectionState<E, P>
where
E: SketchEntry,
{
pub fn new(seed: u64, policy: P) -> Self {
Self {
has_result: false,
table: SketchHashTable::from_raw_parts(
0,
0,
ResizeFactor::X1,
1.0,
MAX_THETA,
seed,
false,
),
policy,
}
}
pub fn update<S>(&mut self, sketch: S) -> Result<(), Error>
where
S: EntrySketch<Entry = E>,
E: Clone,
P: IntersectionMergePolicy<E>,
{
let SketchScalars {
seed_hash,
theta,
empty,
ordered,
num_retained,
} = sketch.scalars();
let new_default_table = |table: &SketchHashTable<E>| {
SketchHashTable::from_raw_parts(
0,
0,
ResizeFactor::X1,
1.0,
table.theta(),
table.seed(),
table.is_empty(),
)
};
if self.table.is_empty() {
return Ok(());
}
if empty {
self.table.set_empty(true);
} else {
check_seed_hash(
self.table.seed_hash(),
seed_hash,
"intersection update",
ErrorKind::InvalidArgument,
)?;
}
self.table.set_theta(if self.table.is_empty() {
MAX_THETA
} else {
self.table.theta().min(theta)
});
if self.has_result && self.table.num_retained() == 0 {
return Ok(());
}
if num_retained == 0 {
self.has_result = true;
self.table = new_default_table(&self.table);
return Ok(());
}
if !self.has_result {
self.has_result = true;
let lg_size = SketchHashTable::<E>::lg_size_from_count_for_rebuild(
num_retained,
HASH_TABLE_REBUILD_THRESHOLD,
);
debug_assert!(lg_size >= 1);
self.table = SketchHashTable::from_raw_parts(
lg_size,
lg_size - 1,
ResizeFactor::X1,
1.0,
self.table.theta(),
self.table.seed(),
self.table.is_empty(),
);
for entry in sketch.entries() {
let hash = entry.hash();
if !self.table.upsert_entry(hash, |existing| match existing {
Some(_) => None,
None => Some(entry),
}) {
return Err(Error::invalid_argument(
"Insert entries from sketch fail, possibly corrupted input sketch",
));
}
}
if self.table.num_retained() != num_retained {
return Err(Error::invalid_argument(
"num entries mismatch, possibly corrupted input sketch",
));
}
} else {
let max_matches = self.table.num_retained().min(num_retained);
let mut matched_entries = Vec::with_capacity(max_matches);
let mut count = 0;
for entry in sketch.entries() {
let hash = entry.hash();
if hash < self.table.theta() {
if let Some(existing) = self.table.entry(hash) {
if matched_entries.len() == max_matches {
return Err(Error::invalid_argument(
"max matches exceeded, possibly corrupted input sketch",
));
}
let mut merged = existing.clone();
self.policy.merge(&mut merged, entry);
matched_entries.push(merged);
}
} else if ordered {
break; }
count += 1;
}
if count > num_retained {
return Err(Error::invalid_argument(
"more keys than expected, possibly corrupted input sketch",
));
} else if !ordered && count < num_retained {
return Err(Error::invalid_argument(
"fewer keys than expected, possibly corrupted input sketch",
));
}
if matched_entries.is_empty() {
self.table = new_default_table(&self.table);
if self.table.theta() == MAX_THETA {
self.table.set_empty(true);
}
} else {
let lg_size = SketchHashTable::<E>::lg_size_from_count_for_rebuild(
matched_entries.len(),
HASH_TABLE_REBUILD_THRESHOLD,
);
debug_assert!(lg_size >= 1);
self.table = SketchHashTable::from_raw_parts(
lg_size,
lg_size - 1,
ResizeFactor::X1,
1.0,
self.table.theta(),
self.table.seed(),
self.table.is_empty(),
);
for entry in matched_entries {
let hash = entry.hash();
if !self.table.upsert_entry(hash, |existing| match existing {
Some(_) => None,
None => Some(entry),
}) {
return Err(Error::invalid_argument(
"duplicate key, possibly corrupted input sketch",
));
}
}
}
}
Ok(())
}
pub fn has_result(&self) -> bool {
self.has_result
}
pub fn estimated_size(&self) -> usize {
self.table.estimated_size()
}
pub fn to_compact_parts(&self, ordered: bool) -> CompactSketchParts<E>
where
E: Clone,
{
let mut entries: Vec<E> = self.table.iter_entries().cloned().collect();
if ordered {
entries.sort_unstable_by_key(SketchEntry::hash);
}
CompactSketchParts {
entries,
theta: self.table.theta(),
seed_hash: self.table.seed_hash(),
ordered,
empty: self.table.is_empty(),
}
}
}