use crate::error::Error;
use crate::error::ErrorKind;
use crate::hash::check_seed_hash;
use crate::hash::compute_seed_hash;
use crate::thetacommon::EntrySketch;
use crate::thetacommon::SketchEntry;
use crate::thetacommon::constants::HASH_TABLE_REBUILD_THRESHOLD;
use crate::thetacommon::constants::MAX_THETA;
use crate::thetacommon::hash_table::SketchHashTable;
use crate::thetacommon::sketch_state::CompactSketchState;
use crate::thetacommon::sketch_state::ThetaFamilySketchMetadata;
pub trait IntersectionMergePolicy<E> {
fn merge(&self, existing: &mut E, incoming: E);
}
#[derive(Debug)]
pub struct IntersectionState<E, P> {
table: SketchHashTable<E>,
policy: P,
result_state: IntersectionResultState,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum IntersectionResultState {
Uninitialized,
Empty,
NonEmpty,
}
impl<E, P> IntersectionState<E, P>
where
E: SketchEntry,
{
pub fn new(seed: u64, policy: P) -> Result<Self, Error> {
let seed_hash = compute_seed_hash(seed, ErrorKind::InvalidArgument)?;
Ok(Self {
result_state: IntersectionResultState::Uninitialized,
table: SketchHashTable::for_set_operation(0, 0, MAX_THETA, seed, seed_hash),
policy,
})
}
pub fn update<S>(&mut self, sketch: S) -> Result<(), Error>
where
S: EntrySketch<Entry = E>,
E: Clone,
P: IntersectionMergePolicy<E>,
{
let table_without_entries = |table: &SketchHashTable<E>, retention_theta| {
SketchHashTable::for_set_operation(
0,
0,
retention_theta,
table.seed(),
table.seed_hash(),
)
};
if self.result_state == IntersectionResultState::Empty {
return Ok(());
}
let (seed_hash, input_theta, input_ordered, input_num_retained) = match sketch.metadata() {
ThetaFamilySketchMetadata::Empty { .. } => {
self.result_state = IntersectionResultState::Empty;
self.table = table_without_entries(&self.table, MAX_THETA);
return Ok(());
}
ThetaFamilySketchMetadata::NonEmpty {
seed_hash,
theta,
ordered,
num_retained,
} => (seed_hash, theta, ordered, num_retained),
};
check_seed_hash(
self.table.seed_hash(),
seed_hash,
"intersection update",
ErrorKind::InvalidArgument,
)?;
let result_theta = self.table.retention_theta().min(input_theta);
self.table.set_retention_theta(result_theta);
if self.result_state == IntersectionResultState::NonEmpty && self.table.num_retained() == 0
{
return Ok(());
}
if input_num_retained == 0 {
self.result_state = IntersectionResultState::NonEmpty;
self.table = table_without_entries(&self.table, result_theta);
return Ok(());
}
if self.result_state == IntersectionResultState::Uninitialized {
self.result_state = IntersectionResultState::NonEmpty;
let lg_size = SketchHashTable::<E>::lg_size_from_count_for_rebuild(
input_num_retained,
HASH_TABLE_REBUILD_THRESHOLD,
);
debug_assert!(lg_size >= 1);
self.table = SketchHashTable::for_set_operation(
lg_size,
lg_size - 1,
result_theta,
self.table.seed(),
self.table.seed_hash(),
);
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() != input_num_retained {
return Err(Error::invalid_argument(
"num entries mismatch, possibly corrupted input sketch",
));
}
} else {
let max_matches = self.table.num_retained().min(input_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.retention_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 input_ordered {
break; }
count += 1;
}
if count > input_num_retained {
return Err(Error::invalid_argument(
"more keys than expected, possibly corrupted input sketch",
));
} else if !input_ordered && count < input_num_retained {
return Err(Error::invalid_argument(
"fewer keys than expected, possibly corrupted input sketch",
));
}
if matched_entries.is_empty() {
self.table = table_without_entries(&self.table, result_theta);
if result_theta == MAX_THETA {
self.result_state = IntersectionResultState::Empty;
}
} 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::for_set_operation(
lg_size,
lg_size - 1,
result_theta,
self.table.seed(),
self.table.seed_hash(),
);
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.result_state != IntersectionResultState::Uninitialized
}
pub fn estimated_size(&self) -> usize {
self.table.estimated_size()
}
pub fn to_compact_sketch_state(&self, ordered: bool) -> Option<CompactSketchState<E>>
where
E: Clone,
{
match self.result_state {
IntersectionResultState::Uninitialized => None,
IntersectionResultState::Empty => {
Some(CompactSketchState::empty(self.table.seed_hash()))
}
IntersectionResultState::NonEmpty => {
Some(self.table.to_non_empty_compact_state(ordered))
}
}
}
}