use std::{
any::{Any, TypeId},
hash::{Hash, Hasher},
ops::Deref,
};
use crate::numeric_id::{DenseIdMap, IdVec, NumericId, define_id};
use crossbeam_queue::SegQueue;
use dashmap::SharedValue;
use rustc_hash::FxHasher;
use crate::{
ColumnId, CounterId, ExecutionState, Offset, SubsetRef, TableId, TaggedRowBuffer, Value,
WrappedTable,
common::{DashMap, IndexSet, SubsetTracker},
parallel,
parallel_heuristics::{parallelize_inter_container_op, parallelize_intra_container_op},
table_spec::{Rebuilder, ValueRebuilder},
};
#[cfg(test)]
mod tests;
define_id!(pub ContainerValueId, u32, "an identifier for containers");
pub trait MergeFn:
Fn(&mut ExecutionState, Value, Value) -> Value + dyn_clone::DynClone + Send + Sync
{
}
impl<T: Fn(&mut ExecutionState, Value, Value) -> Value + Clone + Send + Sync> MergeFn for T {}
dyn_clone::clone_trait_object!(MergeFn);
#[derive(Clone, Default)]
struct ContainerIds {
ids: IndexSet<TypeId>,
}
impl ContainerIds {
fn insert(&mut self, ty: TypeId) -> ContainerValueId {
if let Some(idx) = self.ids.get_index_of(&ty) {
ContainerValueId::from_usize(idx)
} else {
let idx = self.ids.len();
self.ids.insert(ty);
ContainerValueId::from_usize(idx)
}
}
fn get(&self, ty: &TypeId) -> Option<ContainerValueId> {
self.ids.get_index_of(ty).map(ContainerValueId::from_usize)
}
}
#[derive(Clone, Default)]
pub struct ContainerValues {
subset_tracker: SubsetTracker,
container_ids: ContainerIds,
data: DenseIdMap<ContainerValueId, Box<dyn DynamicContainerEnv + Send + Sync>>,
}
#[derive(Clone, Default)]
pub struct ContainerRebuildSummary {
changed: bool,
dirty_ids: IndexSet<Value>,
}
impl ContainerRebuildSummary {
pub fn changed(&self) -> bool {
self.changed
}
pub fn dirty_ids(&self) -> &IndexSet<Value> {
&self.dirty_ids
}
fn note_change(&mut self) {
self.changed = true;
}
fn note_dirty_id(&mut self, value: Value) {
self.changed = true;
self.dirty_ids.insert(value);
}
fn extend(&mut self, other: Self) {
self.changed |= other.changed;
self.dirty_ids.extend(other.dirty_ids);
}
}
impl ContainerValues {
pub fn new() -> Self {
Default::default()
}
fn get<C: ContainerValue>(&self) -> Option<&ContainerEnv<C>> {
let id = self.container_ids.get(&TypeId::of::<C>())?;
let res = self.data.get(id)?.as_any();
Some(res.downcast_ref::<ContainerEnv<C>>().unwrap())
}
pub fn for_each<C: ContainerValue>(&self, mut f: impl FnMut(&C, Value)) {
let Some(env) = self.get::<C>() else {
return;
};
for ent in env.to_id.iter() {
f(ent.key(), *ent.value());
}
}
pub fn get_val<C: ContainerValue>(&self, val: Value) -> Option<impl Deref<Target = C> + '_> {
self.get::<C>()?.get_container(val)
}
pub fn register_val<C: ContainerValue>(
&self,
container: C,
exec_state: &mut ExecutionState,
) -> Value {
let env = self
.get::<C>()
.expect("must register container type before registering a value");
env.get_or_insert(&container, exec_state)
}
pub fn rebuild_val_with(
&self,
type_id: TypeId,
value: Value,
exec_state: &mut ExecutionState,
remap: &(dyn Fn(Value) -> Value + Send + Sync),
) -> Option<Value> {
let id = self.container_ids.get(&type_id)?;
let env = self.data.get(id)?;
env.rebuild_val_with(value, exec_state, remap)
}
pub fn rebuild_all(
&mut self,
table_id: TableId,
table: &WrappedTable,
exec_state: &mut ExecutionState,
) -> ContainerRebuildSummary {
let Some(rebuilder) = table.rebuilder(&[]) else {
return Default::default();
};
let to_scan = rebuilder.hint_col().map(|_| {
self.subset_tracker.recent_updates(table_id, table)
});
let mut summary = if parallelize_inter_container_op(self.data.next_id().index()) {
parallel::map_dense_id_map_mut(&mut self.data, |_, env| {
let mut exec_state = exec_state.clone();
env.apply_rebuild(
table,
&*rebuilder,
to_scan.as_ref().map(|x| x.as_ref()),
&mut exec_state,
)
})
.into_iter()
.fold(ContainerRebuildSummary::default(), |mut acc, summary| {
acc.extend(summary);
acc
})
} else {
let mut summary = ContainerRebuildSummary::default();
for (_, env) in self.data.iter_mut() {
summary.extend(env.apply_rebuild(
table,
&*rebuilder,
to_scan.as_ref().map(|x| x.as_ref()),
exec_state,
));
}
summary
};
self.expand_dirty_id_closure(&mut summary);
summary
}
fn expand_dirty_id_closure(&self, summary: &mut ContainerRebuildSummary) {
let mut frontier = summary.dirty_ids.clone();
let mut seen = frontier.iter().copied().collect::<IndexSet<_>>();
while !frontier.is_empty() {
let mut next = IndexSet::default();
for (_, env) in self.data.iter() {
env.extend_containers_containing(&frontier, &mut next);
}
frontier.clear();
for value in next {
if seen.insert(value) {
summary.note_dirty_id(value);
frontier.insert(value);
}
}
}
}
pub fn register_type<C: ContainerValue>(
&mut self,
id_counter: CounterId,
merge_fn: impl MergeFn + 'static,
) -> ContainerValueId {
let id = self.container_ids.insert(TypeId::of::<C>());
self.data.get_or_insert(id, || {
Box::new(ContainerEnv::<C>::new(Box::new(merge_fn), id_counter))
});
id
}
}
pub trait ContainerValue: Hash + Eq + Clone + Send + Sync + 'static {
fn rebuild_contents(&mut self, rebuilder: &dyn ValueRebuilder) -> bool;
fn iter(&self) -> impl Iterator<Item = Value> + '_;
}
pub trait DynamicContainerEnv: Any + dyn_clone::DynClone + Send + Sync {
fn as_any(&self) -> &dyn Any;
fn apply_rebuild(
&mut self,
table: &WrappedTable,
rebuilder: &dyn Rebuilder,
subset: Option<SubsetRef>,
exec_state: &mut ExecutionState,
) -> ContainerRebuildSummary;
fn extend_containers_containing(&self, values: &IndexSet<Value>, out: &mut IndexSet<Value>);
fn rebuild_val_with(
&self,
value: Value,
exec_state: &mut ExecutionState,
remap: &(dyn Fn(Value) -> Value + Send + Sync),
) -> Option<Value>;
}
dyn_clone::clone_trait_object!(DynamicContainerEnv);
fn hash_container(container: &impl ContainerValue) -> u64 {
let mut hasher = FxHasher::default();
container.hash(&mut hasher);
hasher.finish()
}
#[derive(Clone)]
struct ContainerEnv<C: Eq + Hash> {
merge_fn: Box<dyn MergeFn>,
counter: CounterId,
to_id: DashMap<C, Value>,
to_container: DashMap<Value, (usize /* hash code */, usize /* map */)>,
val_index: DashMap<Value, IndexSet<Value>>,
}
impl<C: ContainerValue> DynamicContainerEnv for ContainerEnv<C> {
fn as_any(&self) -> &dyn Any {
self
}
fn apply_rebuild(
&mut self,
table: &WrappedTable,
rebuilder: &dyn Rebuilder,
subset: Option<SubsetRef>,
exec_state: &mut ExecutionState,
) -> ContainerRebuildSummary {
if let Some(subset) = subset
&& incremental_rebuild(
subset.size(),
self.to_id.len(),
parallelize_intra_container_op(self.to_id.len()),
)
{
return self.apply_rebuild_incremental(
table,
rebuilder,
exec_state,
subset,
rebuilder.hint_col().unwrap(),
);
}
self.apply_rebuild_nonincremental(rebuilder, exec_state)
}
fn extend_containers_containing(&self, values: &IndexSet<Value>, out: &mut IndexSet<Value>) {
for value in values {
if let Some(containers) = self.val_index.get(value) {
out.extend(containers.iter().copied());
}
}
}
fn rebuild_val_with(
&self,
value: Value,
exec_state: &mut ExecutionState,
remap: &(dyn Fn(Value) -> Value + Send + Sync),
) -> Option<Value> {
let mut container = self.get_container(value)?.clone();
container.rebuild_contents(&ClosureRebuilder { remap });
Some(self.get_or_insert(&container, exec_state))
}
}
impl<C: ContainerValue> ContainerEnv<C> {
pub fn new(merge_fn: Box<dyn MergeFn>, counter: CounterId) -> Self {
Self {
merge_fn,
counter,
to_id: DashMap::default(),
to_container: DashMap::default(),
val_index: DashMap::default(),
}
}
fn get_or_insert(&self, container: &C, exec_state: &mut ExecutionState) -> Value {
if let Some(value) = self.to_id.get(container) {
return *value;
}
let value = Value::from_usize(exec_state.inc_counter(self.counter));
let target_map = self.to_id.determine_map(container);
debug_assert_eq!(
target_map,
self.to_container
.determine_shard(hash_container(container) as usize)
);
self.to_container
.insert(value, (hash_container(container) as usize, target_map));
match self.to_id.entry(container.clone()) {
dashmap::Entry::Vacant(vac) => {
vac.insert(value);
for val in container.iter() {
self.val_index.entry(val).or_default().insert(value);
}
value
}
dashmap::Entry::Occupied(occ) => {
let res = *occ.get();
std::mem::drop(occ); self.to_container.remove(&value);
res
}
}
}
fn insert_owned(&self, container: C, value: Value, exec_state: &mut ExecutionState) -> Value {
let hc = hash_container(&container);
let target_map = self.to_id.determine_map(&container);
match self.to_id.entry(container) {
dashmap::Entry::Occupied(mut occ) => {
let result = (self.merge_fn)(exec_state, *occ.get(), value);
let old_val = *occ.get();
if result != old_val {
self.to_container.remove(&old_val);
self.to_container.insert(result, (hc as usize, target_map));
*occ.get_mut() = result;
for val in occ.key().iter() {
let mut index = self.val_index.entry(val).or_default();
index.swap_remove(&old_val);
index.insert(result);
}
}
result
}
dashmap::Entry::Vacant(vacant_entry) => {
self.to_container.insert(value, (hc as usize, target_map));
for val in vacant_entry.key().iter() {
self.val_index.entry(val).or_default().insert(value);
}
vacant_entry.insert(value);
value
}
}
}
fn reinsert_incremental(
&self,
container: C,
old_id: Value,
rebuilt_id: Value,
container_changed: bool,
exec_state: &mut ExecutionState,
summary: &mut ContainerRebuildSummary,
) {
if container_changed || rebuilt_id != old_id {
summary.note_change();
}
if rebuilt_id != old_id {
self.to_container.remove(&old_id);
}
let actual = self.insert_owned(container, rebuilt_id, exec_state);
if container_changed && rebuilt_id == old_id && actual == old_id {
summary.note_dirty_id(old_id);
}
}
fn apply_rebuild_incremental(
&mut self,
table: &WrappedTable,
rebuilder: &dyn Rebuilder,
exec_state: &mut ExecutionState,
to_scan: SubsetRef,
search_col: ColumnId,
) -> ContainerRebuildSummary {
let mut summary = ContainerRebuildSummary::default();
let mut buf = TaggedRowBuffer::new(1);
table.scan_project(
to_scan,
&[search_col],
Offset::new(0),
usize::MAX,
&[],
&mut buf,
);
let mut to_rebuild = IndexSet::<Value>::default();
for (_, row) in buf.iter() {
to_rebuild.insert(row[0]);
let Some(ids) = self.val_index.get(&row[0]) else {
continue;
};
to_rebuild.extend(&*ids);
}
for id in to_rebuild {
let Some((hc, target_map)) = self.to_container.get(&id).map(|x| *x) else {
continue;
};
let shard_mut = self.to_id.shards_mut()[target_map].get_mut();
let Some((mut container, _)) =
shard_mut.remove_entry(hc as u64, |(_, v)| *v.get() == id)
else {
continue;
};
let rebuilt_id = rebuilder.rebuild_val(id);
let container_changed = container.rebuild_contents(rebuilder);
self.reinsert_incremental(
container,
id,
rebuilt_id,
container_changed,
exec_state,
&mut summary,
);
}
summary
}
fn apply_rebuild_nonincremental(
&mut self,
rebuilder: &dyn Rebuilder,
exec_state: &mut ExecutionState,
) -> ContainerRebuildSummary {
if parallelize_inter_container_op(self.to_id.len()) {
return self.apply_rebuild_nonincremental_parallel(rebuilder, exec_state);
}
let mut summary = ContainerRebuildSummary::default();
let mut to_reinsert = Vec::new();
let shards = self.to_id.shards_mut();
for shard in shards.iter_mut() {
let shard = shard.get_mut();
for bucket in unsafe { shard.iter() } {
let (container, val) = unsafe { bucket.as_mut() };
let old_val = *val.get();
let new_val = rebuilder.rebuild_val(old_val);
let container_changed = container.rebuild_contents(rebuilder);
if !container_changed && new_val == old_val {
continue;
}
summary.note_change();
if container_changed {
let ((container, _), _) = unsafe { shard.remove(bucket) };
self.to_container.remove(&old_val);
to_reinsert.push((container, new_val, new_val == old_val));
} else {
*val.get_mut() = new_val;
let prev = self.to_container.remove(&old_val).unwrap().1;
self.to_container.insert(new_val, prev);
}
}
}
for (container, val, stable_id) in to_reinsert {
let actual = self.insert_owned(container, val, exec_state);
if stable_id && actual == val {
summary.note_dirty_id(val);
}
}
summary
}
fn apply_rebuild_nonincremental_parallel(
&mut self,
rebuilder: &dyn Rebuilder,
exec_state: &mut ExecutionState,
) -> ContainerRebuildSummary {
let mut to_reinsert =
IdVec::<usize , SegQueue<(C, Value, bool)>>::default();
to_reinsert.resize_with(self.to_id.shards().len(), Default::default);
let shards = self.to_id.shards_mut();
let changed = parallel::map_mut(shards, |_, shard| {
let mut changed = false;
let shard = shard.get_mut();
for bucket in unsafe { shard.iter() } {
let (container, val) = unsafe { bucket.as_mut() };
let old_val = *val.get();
let new_val = rebuilder.rebuild_val(old_val);
let container_changed = container.rebuild_contents(rebuilder);
if !container_changed && new_val == old_val {
continue;
}
changed = true;
if container_changed {
let ((container, _), _) = unsafe { shard.remove(bucket) };
self.to_container.remove(&old_val);
let shard = self
.to_container
.determine_shard(hash_container(&container) as usize);
to_reinsert[shard].push((container, new_val, new_val == old_val));
} else {
*val.get_mut() = new_val;
let prev = self.to_container.remove(&old_val).unwrap().1;
self.to_container.insert(new_val, prev);
}
}
changed
})
.into_iter()
.any(|changed| changed);
let dirty_ids = SegQueue::new();
parallel::for_each_mut(shards, |shard_id, shard| {
let mut exec_state = exec_state.clone();
let shard = shard.get_mut();
let queue = &to_reinsert[shard_id];
while let Some((container, val, stable_id)) = queue.pop() {
let hc = hash_container(&container);
let target_map = self.to_container.determine_shard(hc as usize);
match shard.find_or_find_insert_slot(
hc,
|(c, _)| c == &container,
|(c, _)| hash_container(c),
) {
Ok(bucket) => {
let (container, val_slot) = unsafe { bucket.as_mut() };
let old_val = *val_slot.get();
let result = (self.merge_fn)(&mut exec_state, old_val, val);
if result != old_val {
self.to_container.remove(&old_val);
self.to_container.insert(result, (hc as usize, target_map));
*val_slot.get_mut() = result;
for val in container.iter() {
let mut index = self.val_index.entry(val).or_default();
index.swap_remove(&old_val);
index.insert(result);
}
}
if stable_id && result == val {
dirty_ids.push(val);
}
}
Err(slot) => {
self.to_container.insert(val, (hc as usize, target_map));
for v in container.iter() {
self.val_index.entry(v).or_default().insert(val);
}
unsafe {
shard.insert_in_slot(hc, slot, (container, SharedValue::new(val)));
}
if stable_id {
dirty_ids.push(val);
}
}
}
}
});
let mut summary = ContainerRebuildSummary::default();
if changed {
summary.note_change();
}
while let Some(value) = dirty_ids.pop() {
summary.note_dirty_id(value);
}
summary
}
fn get_container(&self, value: Value) -> Option<impl Deref<Target = C> + '_> {
let (hc, target_map) = *self.to_container.get(&value)?;
let shard = &self.to_id.shards()[target_map];
let read_guard = shard.read();
let val_ptr: *const (C, _) = shard
.read()
.find(hc as u64, |(_, v)| *v.get() == value)?
.as_ptr();
struct ValueDeref<'a, T, Guard> {
_guard: Guard,
data: &'a T,
}
impl<T, Guard> Deref for ValueDeref<'_, T, Guard> {
type Target = T;
fn deref(&self) -> &T {
self.data
}
}
Some(ValueDeref {
_guard: read_guard,
data: unsafe {
let unwrapped: &(C, _) = &*val_ptr;
&unwrapped.0
},
})
}
}
fn incremental_rebuild(uf_size: usize, table_size: usize, parallel: bool) -> bool {
if parallel {
table_size > 1000 && uf_size * 512 <= table_size
} else {
table_size > 1000 && uf_size * 8 <= table_size
}
}
struct ClosureRebuilder<'a> {
remap: &'a (dyn Fn(Value) -> Value + Send + Sync),
}
impl ValueRebuilder for ClosureRebuilder<'_> {
fn rebuild_val(&self, val: Value) -> Value {
(self.remap)(val)
}
}