use parking_lot::Mutex;
use std::{
collections::HashMap,
hash::Hash,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
};
use super::CellValue;
use crate::{
pipeline::{Definite, Pipeline, PipelineInstall, PipelineSeed},
signal::Signal,
subscription::SubscriptionGuard,
};
type TransitionFn<S, R> = Arc<dyn Fn(&S, &S) -> R + Send + Sync>;
type StateFn<S> = Arc<dyn Fn(&S) + Send + Sync>;
type GuardFn<S> = Arc<dyn Fn(&S, &S) -> bool + Send + Sync>;
type InvalidFn<S> = Arc<dyn Fn(&S, &S) + Send + Sync>;
pub struct StateMachineBuilder<S, R> {
transitions: HashMap<(S, S), TransitionFn<S, R>>,
on_enter: HashMap<S, StateFn<S>>,
on_exit: HashMap<S, StateFn<S>>,
guards: HashMap<(S, S), GuardFn<S>>,
on_any_enter: Vec<StateFn<S>>,
on_invalid: Option<InvalidFn<S>>,
default: Option<R>,
}
impl<S: Eq + Hash + CellValue, R: CellValue> StateMachineBuilder<S, R> {
fn new() -> Self {
Self {
transitions: HashMap::new(),
on_enter: HashMap::new(),
on_exit: HashMap::new(),
guards: HashMap::new(),
on_any_enter: Vec::new(),
on_invalid: None,
default: None,
}
}
pub fn with_default(&mut self, value: R) -> &mut Self {
self.default = Some(value);
self
}
pub fn on<F>(&mut self, from: S, to: S, handler: F) -> &mut Self
where
F: Fn(&S, &S) -> R + Send + Sync + 'static,
{
self.transitions.insert((from, to), Arc::new(handler));
self
}
pub fn on_enter<F>(&mut self, state: S, handler: F) -> &mut Self
where
F: Fn(&S) + Send + Sync + 'static,
{
self.on_enter.insert(state, Arc::new(handler));
self
}
pub fn on_exit<F>(&mut self, state: S, handler: F) -> &mut Self
where
F: Fn(&S) + Send + Sync + 'static,
{
self.on_exit.insert(state, Arc::new(handler));
self
}
pub fn on_any<F>(&mut self, handler: F) -> &mut Self
where
F: Fn(&S) + Send + Sync + 'static,
{
self.on_any_enter.push(Arc::new(handler));
self
}
pub fn guard<F>(&mut self, from: S, to: S, predicate: F) -> &mut Self
where
F: Fn(&S, &S) -> bool + Send + Sync + 'static,
{
self.guards.insert((from, to), Arc::new(predicate));
self
}
pub fn on_invalid<F>(&mut self, handler: F) -> &mut Self
where
F: Fn(&S, &S) + Send + Sync + 'static,
{
self.on_invalid = Some(Arc::new(handler));
self
}
}
pub struct StateTransitionPipeline<P, S, R> {
source: P,
transitions: Arc<HashMap<(S, S), TransitionFn<S, R>>>,
on_enter: Arc<HashMap<S, StateFn<S>>>,
on_exit: Arc<HashMap<S, StateFn<S>>>,
guards: Arc<HashMap<(S, S), GuardFn<S>>>,
on_any_enter: Arc<Vec<StateFn<S>>>,
on_invalid: Option<InvalidFn<S>>,
initial: R,
}
impl<P, S, R> PipelineInstall<R> for StateTransitionPipeline<P, S, R>
where
P: PipelineInstall<S> + PipelineSeed<S>,
S: CellValue + Eq + Hash,
R: CellValue,
{
fn install(&self, callback: Arc<dyn Fn(&Signal<R>) + Send + Sync>) -> SubscriptionGuard {
let transitions = self.transitions.clone();
let on_enter = self.on_enter.clone();
let on_exit = self.on_exit.clone();
let guards = self.guards.clone();
let on_any_enter = self.on_any_enter.clone();
let on_invalid = self.on_invalid.clone();
let first = AtomicBool::new(true);
let current_state = Arc::new(Mutex::new(self.source.seed()));
self.source.install(Arc::new(move |signal| match signal {
Signal::Value(next) => {
if first.swap(false, Ordering::SeqCst) {
return;
}
let current = {
let mut guard = current_state.lock();
let previous = guard.clone();
*guard = next.as_ref().clone();
previous
};
let key = (current.clone(), next.as_ref().clone());
if !transitions.contains_key(&key) {
if let Some(handler) = &on_invalid {
handler(¤t, next);
}
return;
}
if let Some(guard) = guards.get(&key)
&& !guard(¤t, next)
{
return;
}
if let Some(handler) = on_exit.get(¤t) {
handler(¤t);
}
let output = transitions.get(&key).map(|handler| handler(¤t, next));
if let Some(handler) = on_enter.get(next.as_ref()) {
handler(next);
}
for handler in on_any_enter.iter() {
handler(next);
}
if let Some(value) = output {
callback(&Signal::value(value));
}
}
Signal::Complete => callback(&Signal::Complete),
Signal::Error(error) => callback(&Signal::Error(error.clone())),
}))
}
}
impl<P, S, R> PipelineSeed<R> for StateTransitionPipeline<P, S, R>
where
P: PipelineInstall<S> + PipelineSeed<S>,
S: CellValue + Eq + Hash,
R: CellValue,
{
fn seed(&self) -> R {
self.initial.clone()
}
}
#[allow(private_bounds)]
impl<P, S, R> Pipeline<R, Definite> for StateTransitionPipeline<P, S, R>
where
P: Pipeline<S, Definite> + PipelineSeed<S>,
S: CellValue + Eq + Hash,
R: CellValue,
{
}
#[allow(private_bounds)]
pub trait StateTransitionExt<S: CellValue + Eq + Hash>:
Pipeline<S, Definite> + PipelineSeed<S>
{
#[track_caller]
fn state_transition<R, F>(self, configure: F) -> impl crate::Materialize<R, Definite>
where
S: CellValue + Eq + Hash,
R: CellValue + Default,
F: FnOnce(&mut StateMachineBuilder<S, R>),
{
let mut builder = StateMachineBuilder::new();
configure(&mut builder);
let initial = builder.default.take().unwrap_or_default();
StateTransitionPipeline {
source: self,
transitions: Arc::new(builder.transitions),
on_enter: Arc::new(builder.on_enter),
on_exit: Arc::new(builder.on_exit),
guards: Arc::new(builder.guards),
on_any_enter: Arc::new(builder.on_any_enter),
on_invalid: builder.on_invalid,
initial,
}
}
}
impl<S, P> StateTransitionExt<S> for P
where
S: CellValue + Eq + Hash,
P: Pipeline<S, Definite> + PipelineSeed<S>,
{
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicU32, Ordering};
use super::*;
use crate::{Cell, Materialize, Mutable, traits::Watchable};
#[derive(Clone, PartialEq, Eq, Hash, Debug)]
enum State {
Idle,
Loading,
Ready,
Error,
}
#[test]
fn test_state_transition_valid() {
let source = Cell::new(State::Idle);
let transition_count = Arc::new(AtomicU32::new(0));
let tc = transition_count.clone();
let sm = source
.clone()
.state_transition(|sm| {
sm.on(State::Idle, State::Loading, move |_, _| {
tc.fetch_add(1, Ordering::SeqCst);
true
});
sm.on(State::Loading, State::Ready, |_, _| true);
sm.on(State::Loading, State::Error, |_, _| true);
})
.materialize();
let emissions = Arc::new(AtomicU32::new(0));
let e = emissions.clone();
let _guard = sm.subscribe(move |_| {
e.fetch_add(1, Ordering::SeqCst);
});
assert_eq!(emissions.load(Ordering::SeqCst), 1);
source.set(State::Loading);
assert_eq!(emissions.load(Ordering::SeqCst), 2);
assert_eq!(transition_count.load(Ordering::SeqCst), 1);
source.set(State::Ready);
assert_eq!(emissions.load(Ordering::SeqCst), 3);
}
#[test]
fn test_state_transition_undefined_advances_state() {
let source = Cell::new(State::Idle);
let sm = source
.clone()
.state_transition(|sm| {
sm.on(State::Idle, State::Loading, |_, _| true);
sm.on(State::Loading, State::Ready, |_, _| true);
})
.materialize();
let emissions = Arc::new(AtomicU32::new(0));
let e = emissions.clone();
let _guard = sm.subscribe(move |_| {
e.fetch_add(1, Ordering::SeqCst);
});
assert_eq!(emissions.load(Ordering::SeqCst), 1);
source.set(State::Ready);
assert_eq!(emissions.load(Ordering::SeqCst), 1);
source.set(State::Error);
assert_eq!(emissions.load(Ordering::SeqCst), 1);
source.set(State::Loading);
assert_eq!(emissions.load(Ordering::SeqCst), 1);
source.set(State::Ready);
assert_eq!(emissions.load(Ordering::SeqCst), 2);
}
#[test]
fn test_state_transition_on_enter_exit() {
let source = Cell::new(State::Idle);
let enter_count = Arc::new(AtomicU32::new(0));
let exit_count = Arc::new(AtomicU32::new(0));
let ec = enter_count.clone();
let xc = exit_count.clone();
let _sm: Cell<bool, _> = source
.clone()
.state_transition(|sm| {
sm.on(State::Idle, State::Loading, |_, _| true);
sm.on_exit(State::Idle, move |_| {
xc.fetch_add(1, Ordering::SeqCst);
});
sm.on_enter(State::Loading, move |_| {
ec.fetch_add(1, Ordering::SeqCst);
});
})
.materialize();
source.set(State::Loading);
assert_eq!(exit_count.load(Ordering::SeqCst), 1);
assert_eq!(enter_count.load(Ordering::SeqCst), 1);
}
#[test]
fn test_state_transition_guard() {
let source = Cell::new(State::Idle);
let allow = Arc::new(AtomicBool::new(false));
let a = allow.clone();
let sm = source
.clone()
.state_transition(|sm| {
sm.on(State::Idle, State::Loading, |_, _| true);
sm.guard(State::Idle, State::Loading, move |_, _| {
a.load(Ordering::SeqCst)
});
})
.materialize();
let emissions = Arc::new(AtomicU32::new(0));
let e = emissions.clone();
let _guard = sm.subscribe(move |_| {
e.fetch_add(1, Ordering::SeqCst);
});
source.set(State::Loading);
assert_eq!(emissions.load(Ordering::SeqCst), 1);
source.set(State::Idle); allow.store(true, Ordering::SeqCst);
let source2 = Cell::new(State::Idle);
let a2 = allow;
let sm2 = source2
.clone()
.state_transition(|sm| {
sm.on(State::Idle, State::Loading, |_, _| true);
sm.guard(State::Idle, State::Loading, move |_, _| {
a2.load(Ordering::SeqCst)
});
})
.materialize();
let emissions2 = Arc::new(AtomicU32::new(0));
let e2 = emissions2.clone();
let _guard2 = sm2.subscribe(move |_| {
e2.fetch_add(1, Ordering::SeqCst);
});
source2.set(State::Loading);
assert_eq!(emissions2.load(Ordering::SeqCst), 2); }
#[test]
fn test_state_transition_on_invalid() {
let source = Cell::new(State::Idle);
let invalid_count = Arc::new(AtomicU32::new(0));
let ic = invalid_count.clone();
let _sm: Cell<bool, _> = source
.clone()
.state_transition(|sm| {
sm.on(State::Idle, State::Loading, |_, _| true);
sm.on_invalid(move |_, _| {
ic.fetch_add(1, Ordering::SeqCst);
});
})
.materialize();
source.set(State::Ready);
assert_eq!(invalid_count.load(Ordering::SeqCst), 1);
source.set(State::Error);
assert_eq!(invalid_count.load(Ordering::SeqCst), 2);
}
#[test]
fn test_state_transition_selective_emit() {
use crate::{FilterExt, Gettable, Materialize};
let source = Cell::new(State::Idle);
let sm = source.clone().state_transition(|sm| {
sm.on(State::Idle, State::Loading, |_, _| true);
sm.on(State::Loading, State::Ready, |_, _| false);
sm.on(State::Ready, State::Idle, |_, _| false);
});
let triggers = sm.filter(|v| *v).materialize();
let emission_count = Arc::new(AtomicU32::new(0));
let ec = emission_count.clone();
let _guard = triggers.subscribe(move |_| {
ec.fetch_add(1, Ordering::SeqCst);
});
assert_eq!(emission_count.load(Ordering::SeqCst), 1);
source.set(State::Loading); assert_eq!(emission_count.load(Ordering::SeqCst), 2);
source.set(State::Ready); assert_eq!(emission_count.load(Ordering::SeqCst), 2);
source.set(State::Idle); assert_eq!(emission_count.load(Ordering::SeqCst), 2);
source.set(State::Loading); assert_eq!(emission_count.load(Ordering::SeqCst), 3);
assert_eq!(triggers.get(), Some(true));
}
}