use std::cell::Cell;
use std::collections::BTreeMap;
use std::sync::{Arc, Mutex, PoisonError};
use crate::shard::commit_state::ShardCommitState;
use crate::tree::Hash;
use super::Database;
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct RootAdvance {
pub shard_id: usize,
#[serde(with = "hash_serde")]
pub prior_root: Hash,
#[serde(with = "hash_serde")]
pub new_root: Hash,
pub advance_gen: u64,
}
#[derive(Debug, Clone, Copy)]
pub struct RootTransition {
pub prior_root: Hash,
pub new_root: Hash,
pub advance_gen: u64,
}
#[derive(Debug)]
pub struct ShardEmitState {
emission: Mutex<u64>,
}
impl ShardEmitState {
const fn new() -> Self {
Self {
emission: Mutex::new(0),
}
}
}
struct Subscriber {
id: u64,
callback: Arc<dyn Fn(RootAdvance) + Send + Sync + 'static>,
}
#[derive(Default)]
struct Registry {
next_id: u64,
subscribers: Vec<Subscriber>,
}
#[derive(Default)]
pub struct RootAdvanceSeam {
registry: Mutex<Registry>,
shards: Mutex<BTreeMap<usize, Arc<ShardEmitState>>>,
commit_states: Mutex<BTreeMap<usize, Arc<ShardCommitState>>>,
}
impl std::fmt::Debug for RootAdvanceSeam {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("RootAdvanceSeam")
.finish_non_exhaustive()
}
}
impl RootAdvanceSeam {
pub fn new() -> Arc<Self> {
Arc::new(Self::default())
}
pub fn shard_state(&self, shard_id: usize) -> Arc<ShardEmitState> {
let mut shards = lock_adopt(&self.shards);
Arc::clone(
shards
.entry(shard_id)
.or_insert_with(|| Arc::new(ShardEmitState::new())),
)
}
pub fn commit_state(&self, shard_id: usize) -> Arc<ShardCommitState> {
let mut cells = lock_adopt(&self.commit_states);
Arc::clone(cells.entry(shard_id).or_insert_with(ShardCommitState::new))
}
fn subscribe(&self, callback: Arc<dyn Fn(RootAdvance) + Send + Sync + 'static>) -> u64 {
let mut registry = lock_adopt(&self.registry);
let id = registry.next_id;
registry.next_id = registry.next_id.wrapping_add(1);
registry.subscribers.push(Subscriber { id, callback });
id
}
fn cancel(&self, id: u64) {
lock_adopt(&self.registry)
.subscribers
.retain(|subscriber| subscriber.id != id);
}
fn snapshot(&self) -> Vec<Arc<dyn Fn(RootAdvance) + Send + Sync + 'static>> {
lock_adopt(&self.registry)
.subscribers
.iter()
.map(|subscriber| Arc::clone(&subscriber.callback))
.collect()
}
pub fn emit(&self, shard_id: usize, state: &ShardEmitState, transition: RootTransition) {
let mut last_told = lock_adopt(&state.emission);
if transition.advance_gen <= *last_told {
return;
}
*last_told = transition.advance_gen;
let subscribers = self.snapshot();
if subscribers.is_empty() {
return;
}
let event = RootAdvance {
shard_id,
prior_root: transition.prior_root,
new_root: transition.new_root,
advance_gen: transition.advance_gen,
};
let _wall = InEmissionGuard::enter();
for callback in &subscribers {
callback(event);
}
drop(last_told);
}
}
fn lock_adopt<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
mutex.lock().unwrap_or_else(PoisonError::into_inner)
}
thread_local! {
static IN_EMISSION: Cell<bool> = const { Cell::new(false) };
}
pub fn in_emission() -> bool {
IN_EMISSION.with(Cell::get)
}
struct InEmissionGuard {
previous: bool,
}
impl InEmissionGuard {
fn enter() -> Self {
let previous = IN_EMISSION.with(|flag| flag.replace(true));
Self { previous }
}
}
impl Drop for InEmissionGuard {
fn drop(&mut self) {
IN_EMISSION.with(|flag| flag.set(self.previous));
}
}
#[must_use = "dropping the subscription immediately unsubscribes; keep it alive to keep receiving tells"]
pub struct RootAdvanceSubscription {
seam: Arc<RootAdvanceSeam>,
id: u64,
active: bool,
}
impl RootAdvanceSubscription {
pub fn cancel(mut self) {
self.deactivate();
}
fn deactivate(&mut self) {
if self.active {
self.active = false;
self.seam.cancel(self.id);
}
}
}
impl std::fmt::Debug for RootAdvanceSubscription {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("RootAdvanceSubscription")
.field("id", &self.id)
.field("active", &self.active)
.finish_non_exhaustive()
}
}
impl Drop for RootAdvanceSubscription {
fn drop(&mut self) {
self.deactivate();
}
}
impl Database {
pub fn subscribe_root_advance(
&self,
callback: impl Fn(RootAdvance) + Send + Sync + 'static,
) -> RootAdvanceSubscription {
let seam = self.seam();
let id = seam.subscribe(Arc::new(callback));
RootAdvanceSubscription {
seam: Arc::clone(seam),
id,
active: true,
}
}
}
pub const WRITE_DURING_EMISSION_REMEDY: &str = "a root-advance subscriber callback attempted an engine write on the same \
Database; write-back from a callback is refused (it would recurse \
commit->tell->commit and self-deadlock on the shard's emission mutex). \
Remedy: record the tell and hand the write to your own executor/queue; \
reads (get/range/diff/checkout) are permitted inside a callback";
#[cfg(test)]
#[path = "root_advance_tests.rs"]
mod tests;
mod hash_serde {
use crate::tree::Hash;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub(super) fn serialize<S: Serializer>(hash: &Hash, serializer: S) -> Result<S::Ok, S::Error> {
hash.as_bytes().serialize(serializer)
}
pub(super) fn deserialize<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Hash, D::Error> {
let bytes = <[u8; crate::tree::node::HASH_SIZE]>::deserialize(deserializer)?;
Ok(Hash::from_bytes(bytes))
}
}