use crate::compat::HashMap;
use crate::compat::{Mutex, RwLock};
use alloc::sync::Arc;
use core::sync::atomic::{AtomicU64, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum Priority {
High,
#[default]
Normal,
Low,
}
impl Priority {
fn rank(&self) -> u8 {
match self {
Priority::High => 0,
Priority::Normal => 1,
Priority::Low => 2,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct ConnectionHandle(pub u64);
static NEXT_HANDLE: AtomicU64 = AtomicU64::new(1);
type SlotFn<T> = Box<dyn FnMut(Arc<T>) + Send + Sync + 'static>;
struct SlotEntry<T: Clone + Send + 'static> {
callback: Option<SlotFn<T>>,
once: bool,
blocked: bool,
priority: Priority,
}
struct SignalInner<T: Clone + Send + 'static> {
slots: RwLock<HashMap<ConnectionHandle, SlotEntry<T>>>,
}
impl<T: Clone + Send + 'static> SignalInner<T> {
fn disconnect(&self, handle: ConnectionHandle) -> bool {
self.slots.write().expect("signal lock poisoned").remove(&handle).is_some()
}
fn block(&self, handle: ConnectionHandle) -> bool {
if let Some(entry) = self.slots.write().expect("signal lock poisoned").get_mut(&handle) {
entry.blocked = true;
true
} else {
false
}
}
fn unblock(&self, handle: ConnectionHandle) -> bool {
if let Some(entry) = self.slots.write().expect("signal lock poisoned").get_mut(&handle) {
entry.blocked = false;
true
} else {
false
}
}
fn is_blocked(&self, handle: ConnectionHandle) -> Option<bool> {
self.slots.read().expect("signal lock poisoned").get(&handle).map(|entry| entry.blocked)
}
fn set_priority(&self, handle: ConnectionHandle, priority: Priority) -> bool {
if let Some(entry) = self.slots.write().expect("signal lock poisoned").get_mut(&handle) {
entry.priority = priority;
true
} else {
false
}
}
}
#[derive(Default)]
pub struct ConnectionScope {
disconnectors: Mutex<Vec<Box<dyn FnOnce() + Send + 'static>>>,
}
impl ConnectionScope {
pub fn new() -> Self {
Self::default()
}
pub fn clear(&self) {
let mut disconnectors = self.disconnectors.lock().unwrap_or_else(|e| e.into_inner());
while let Some(disconnector) = disconnectors.pop() {
disconnector();
}
}
pub fn disconnect_count(&self) -> usize {
self.disconnectors.lock().unwrap_or_else(|e| e.into_inner()).len()
}
fn track(&self, disconnector: Box<dyn FnOnce() + Send + 'static>) {
self.disconnectors.lock().unwrap_or_else(|e| e.into_inner()).push(disconnector);
}
}
impl Drop for ConnectionScope {
fn drop(&mut self) {
let mut disconnectors = self.disconnectors.lock().unwrap_or_else(|e| e.into_inner());
while let Some(disconnector) = disconnectors.pop() {
disconnector();
}
}
}
#[derive(Clone)]
pub struct Signal<T: Clone + Send + 'static> {
inner: Arc<SignalInner<T>>,
}
impl<T: Clone + Send + 'static> Signal<T> {
pub fn new() -> Self {
Self { inner: Arc::new(SignalInner { slots: RwLock::new(HashMap::new()) }) }
}
pub fn connect<F>(&self, slot: F) -> ConnectionHandle
where
F: FnMut(Arc<T>) + Send + Sync + 'static,
{
self.connect_with_priority(slot, Priority::Normal)
}
pub fn connect_with_priority<F>(&self, slot: F, priority: Priority) -> ConnectionHandle
where
F: FnMut(Arc<T>) + Send + Sync + 'static,
{
let handle = ConnectionHandle(NEXT_HANDLE.fetch_add(1, Ordering::Relaxed));
self.inner.slots.write().expect("signal lock poisoned").insert(
handle,
SlotEntry { callback: Some(Box::new(slot)), once: false, blocked: false, priority },
);
handle
}
pub fn connect_once<F>(&self, slot: F) -> ConnectionHandle
where
F: FnMut(Arc<T>) + Send + Sync + 'static,
{
let handle = ConnectionHandle(NEXT_HANDLE.fetch_add(1, Ordering::Relaxed));
self.inner.slots.write().expect("signal lock poisoned").insert(
handle,
SlotEntry {
callback: Some(Box::new(slot)),
once: true,
blocked: false,
priority: Priority::Normal,
},
);
handle
}
pub fn connect_scoped<F>(&self, owner: &ConnectionScope, slot: F) -> ConnectionHandle
where
F: FnMut(Arc<T>) + Send + Sync + 'static,
{
let handle = self.connect(slot);
self.track_owner(owner, handle);
handle
}
pub fn connect_once_scoped<F>(&self, owner: &ConnectionScope, slot: F) -> ConnectionHandle
where
F: FnMut(Arc<T>) + Send + Sync + 'static,
{
let handle = self.connect_once(slot);
self.track_owner(owner, handle);
handle
}
pub fn disconnect(&self, handle: ConnectionHandle) -> bool {
self.inner.disconnect(handle)
}
pub fn disconnect_all(&self) {
self.inner.slots.write().expect("signal lock poisoned").clear();
}
pub fn block(&self, handle: ConnectionHandle) -> bool {
self.inner.block(handle)
}
pub fn unblock(&self, handle: ConnectionHandle) -> bool {
self.inner.unblock(handle)
}
pub fn is_blocked(&self, handle: ConnectionHandle) -> Option<bool> {
self.inner.is_blocked(handle)
}
pub fn is_connected(&self, handle: ConnectionHandle) -> bool {
self.inner.slots.read().expect("signal lock poisoned").contains_key(&handle)
}
pub fn set_priority(&self, handle: ConnectionHandle, priority: Priority) -> bool {
self.inner.set_priority(handle, priority)
}
pub fn emit(&self, value: T) {
let arc_value = Arc::new(value);
let snapshot: Vec<(ConnectionHandle, Priority)> = {
let slots = self.inner.slots.read().expect("signal lock poisoned");
slots.iter().map(|(h, e)| (*h, e.priority)).collect()
};
let mut snapshot = snapshot;
snapshot.sort_by_key(|a| a.1.rank());
for (handle, _priority) in snapshot {
let taken = {
let mut slots = self.inner.slots.write().expect("signal lock poisoned");
if let Some(entry) = slots.get_mut(&handle) {
if entry.blocked {
None
} else {
entry.callback.take()
}
} else {
None
}
};
if let Some(mut callback) = taken {
callback(arc_value.clone());
let mut slots = self.inner.slots.write().expect("signal lock poisoned");
if let Some(entry) = slots.get_mut(&handle) {
if entry.once {
slots.remove(&handle);
} else {
entry.callback = Some(callback);
}
}
}
}
}
pub fn slot_count(&self) -> usize {
self.inner.slots.read().expect("signal lock poisoned").len()
}
fn track_owner(&self, owner: &ConnectionScope, handle: ConnectionHandle) {
let weak = Arc::downgrade(&self.inner);
owner.track(Box::new(move || {
if let Some(inner) = weak.upgrade() {
let _ = inner.disconnect(handle);
}
}));
}
}
impl<T: Clone + Send + 'static> Default for Signal<T> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod emit_behaviour_tests {
use super::*;
use alloc::sync::Arc;
use core::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Mutex;
#[derive(Default)]
struct Trace {
entries: Mutex<alloc::vec::Vec<&'static str>>,
}
impl Trace {
fn push(&self, label: &'static str) {
self.entries.lock().unwrap().push(label);
}
fn snapshot(&self) -> alloc::vec::Vec<&'static str> {
self.entries.lock().unwrap().clone()
}
}
#[test]
fn a_slot_may_disconnect_itself_from_inside_its_callback() {
use core::sync::atomic::AtomicU64;
let signal = Signal::<u32>::new();
let calls = Arc::new(AtomicUsize::new(0));
let shared_handle = Arc::new(AtomicU64::new(0));
let calls_self = Arc::clone(&calls);
let signal_for_self = signal.clone();
let handle_slot = Arc::clone(&shared_handle);
let handle = signal.connect(move |_| {
calls_self.fetch_add(1, Ordering::SeqCst);
let own = ConnectionHandle(handle_slot.load(Ordering::SeqCst));
assert!(
signal_for_self.disconnect(own),
"a callback must be able to find and remove its own handle"
);
});
shared_handle.store(handle.0, Ordering::SeqCst);
signal.emit(1);
assert_eq!(calls.load(Ordering::SeqCst), 1, "the slot must run exactly once");
assert!(
!signal.is_connected(handle),
"self-disconnect must survive the restore step, not be undone by it"
);
signal.emit(2);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"a self-disconnected slot must not be called by the next emit"
);
}
#[test]
fn a_slot_may_disconnect_a_later_slot_in_the_same_pass() {
let signal = Signal::<u32>::new();
let trace = Arc::new(Trace::default());
let target = {
let trace = Arc::clone(&trace);
signal.connect(move |_| trace.push("target"))
};
let signal_for_cutter = signal.clone();
let trace_first = Arc::clone(&trace);
signal.connect_with_priority(
move |_| {
trace_first.push("cutter");
signal_for_cutter.disconnect(target);
},
Priority::High,
);
signal.emit(1);
assert_eq!(
trace.snapshot(),
alloc::vec!["cutter"],
"the disconnected slot must be skipped in the same pass"
);
}
#[test]
fn connecting_inside_a_callback_does_not_run_in_the_same_pass() {
let signal = Signal::<u32>::new();
let trace = Arc::new(Trace::default());
let signal_for_adder = signal.clone();
let trace_adder = Arc::clone(&trace);
signal.connect(move |_| {
trace_adder.push("adder");
let trace_late = Arc::clone(&trace_adder);
signal_for_adder.connect(move |_| trace_late.push("late"));
});
signal.emit(1);
assert_eq!(
trace.snapshot(),
alloc::vec!["adder"],
"a slot connected during emit must wait for the next emit"
);
trace.entries.lock().unwrap().clear();
signal.emit(2);
let mut order = trace.snapshot();
order.sort_unstable();
assert_eq!(
order,
alloc::vec!["adder", "late"],
"both the original and the newly connected slot must run on the next emit"
);
assert_eq!(signal.slot_count(), 3);
}
#[test]
fn a_re_entrant_emit_skips_the_slot_still_on_the_stack() {
let signal = Signal::<u32>::new();
let trace = Arc::new(Trace::default());
let nested_seen = Arc::new(AtomicUsize::new(0));
{
let trace = Arc::clone(&trace);
let nested = Arc::clone(&nested_seen);
signal.connect_with_priority(
move |_| {
trace.push("observer");
nested.fetch_add(1, Ordering::SeqCst);
},
Priority::Low,
);
}
let signal_for_reentry = signal.clone();
let trace_recur = Arc::clone(&trace);
let depth = Arc::new(AtomicUsize::new(0));
let depth_inner = Arc::clone(&depth);
signal.connect_with_priority(
move |_| {
trace_recur.push("recur");
if depth_inner.fetch_add(1, Ordering::SeqCst) == 0 {
signal_for_reentry.emit(0);
}
},
Priority::High,
);
signal.emit(1);
let order = trace.snapshot();
assert_eq!(
order.iter().filter(|l| **l == "recur").count(),
1,
"the recursive slot must run once, not twice: {order:?}"
);
assert!(
order.iter().filter(|l| **l == "observer").count() >= 2,
"the re-entrant pass must still reach the other slot: {order:?}"
);
assert_eq!(
nested_seen.load(Ordering::SeqCst),
2,
"the observer must have been called by both the outer and the nested pass"
);
}
#[test]
fn a_once_slot_runs_once_and_is_removed() {
let signal = Signal::<u32>::new();
let calls = Arc::new(AtomicUsize::new(0));
let calls_once = Arc::clone(&calls);
let handle = signal.connect_once(move |_| {
calls_once.fetch_add(1, Ordering::SeqCst);
});
signal.emit(1);
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert!(!signal.is_connected(handle), "a once slot must be removed after it runs");
signal.emit(2);
assert_eq!(calls.load(Ordering::SeqCst), 1, "a once slot must not run twice");
}
#[test]
fn a_blocked_slot_is_skipped_and_unblocking_restores_it() {
let signal = Signal::<u32>::new();
let calls = Arc::new(AtomicUsize::new(0));
let calls_inner = Arc::clone(&calls);
let handle = signal.connect(move |_| {
calls_inner.fetch_add(1, Ordering::SeqCst);
});
assert!(signal.block(handle));
signal.emit(1);
assert_eq!(calls.load(Ordering::SeqCst), 0, "a blocked slot must not run");
assert!(signal.is_connected(handle), "blocking must not disconnect");
assert!(signal.unblock(handle));
signal.emit(2);
assert_eq!(calls.load(Ordering::SeqCst), 1, "an unblocked slot must run again");
}
#[test]
fn slots_run_in_priority_order() {
let signal = Signal::<u32>::new();
let trace = Arc::new(Trace::default());
for (label, priority) in [
("low", Priority::Low),
("normal", Priority::Normal),
("high", Priority::High),
("normal2", Priority::Normal),
] {
let trace = Arc::clone(&trace);
signal.connect_with_priority(move |_| trace.push(label), priority);
}
signal.emit(1);
let order = trace.snapshot();
assert_eq!(order.first(), Some(&"high"), "High must run first");
assert_eq!(order.last(), Some(&"low"), "Low must run last");
assert_eq!(order.len(), 4, "every non-blocked slot must run exactly once, got {order:?}");
}
#[test]
fn set_priority_inside_a_callback_does_not_reorder_the_current_pass() {
let signal = Signal::<u32>::new();
let trace = Arc::new(Trace::default());
let later = {
let trace = Arc::clone(&trace);
signal.connect_with_priority(move |_| trace.push("later"), Priority::Low)
};
let signal_for_bump = signal.clone();
let trace_first = Arc::clone(&trace);
signal.connect_with_priority(
move |_| {
trace_first.push("first");
signal_for_bump.set_priority(later, Priority::High);
},
Priority::High,
);
signal.emit(1);
assert_eq!(
trace.snapshot(),
alloc::vec!["first", "later"],
"the pass order is fixed at snapshot time"
);
}
#[test]
fn disconnect_all_inside_a_callback_stops_the_rest_of_the_pass() {
let signal = Signal::<u32>::new();
let trace = Arc::new(Trace::default());
{
let trace = Arc::clone(&trace);
signal.connect_with_priority(move |_| trace.push("second"), Priority::Normal);
}
let signal_for_clear = signal.clone();
let trace_first = Arc::clone(&trace);
signal.connect_with_priority(
move |_| {
trace_first.push("first");
signal_for_clear.disconnect_all();
},
Priority::High,
);
signal.emit(1);
assert_eq!(
trace.snapshot(),
alloc::vec!["first"],
"slots after disconnect_all must be skipped"
);
assert_eq!(signal.slot_count(), 0u64 as usize, "no slots must remain");
}
#[test]
fn a_scoped_connection_is_dropped_with_its_scope() {
let signal = Signal::<u32>::new();
let calls = Arc::new(AtomicUsize::new(0));
{
let scope = ConnectionScope::new();
let calls_inner = Arc::clone(&calls);
signal.connect_scoped(&scope, move |_| {
calls_inner.fetch_add(1, Ordering::SeqCst);
});
signal.emit(1);
assert_eq!(calls.load(Ordering::SeqCst), 1, "the scoped slot must run while alive");
}
signal.emit(2);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"the scoped slot must be gone once its scope drops"
);
}
}