use async_trait::async_trait;
use log::info;
use std::collections::HashMap;
use std::fmt::Debug;
use std::hash::Hash;
use tokio::sync::broadcast;
use tokio::sync::mpsc;
use tokio::sync::mpsc::{Receiver, Sender};
use tokio::time::Instant;
pub struct Data<Event, State, UserData> {
pub prev_state: Option<State>,
pub state: State,
pub user_data: UserData,
#[allow(dead_code)]
pub events: HashMap<Event, Instant>,
}
#[async_trait]
pub trait Transition<
Event: Debug + Copy + Clone + PartialEq + Eq + Hash,
State: Default + Debug + Eq + PartialEq + Copy + Clone + Hash,
UserData: Debug + Default,
>
{
async fn next(&mut self, event: Event, data: &Data<Event, State, UserData>) -> State;
fn enter(&mut self, _data: &Data<Event, State, UserData>) {}
}
type FnOnEventRegister<Event, State, UserData> = fn(Event, &mut Data<Event, State, UserData>);
pub struct StateMachine<Event, State, UserData> {
event_receiver: Receiver<Event>,
broadcast: (
tokio::sync::broadcast::Sender<State>,
tokio::sync::broadcast::Receiver<State>,
),
transitions: HashMap<State, Box<dyn Transition<Event, State, UserData> + Send + Sync>>,
data: Data<Event, State, UserData>,
on_event_register: Option<FnOnEventRegister<Event, State, UserData>>,
}
impl<Event, State, UserData> StateMachine<Event, State, UserData>
where
Event: Debug + Copy + Clone + PartialEq + Eq + Hash,
State: Default + Debug + Eq + PartialEq + Copy + Clone + Hash,
UserData: Debug + Default,
{
pub fn new(size: usize) -> (Self, Sender<Event>) {
let (event_sender, event_receiver) = mpsc::channel::<Event>(size);
let fsm = Self {
event_receiver,
broadcast: broadcast::channel::<State>(size),
transitions: HashMap::new(),
data: Data {
prev_state: None,
state: State::default(),
user_data: UserData::default(),
events: HashMap::new(),
},
on_event_register: None,
};
(fsm, event_sender)
}
pub fn add_transition(
&mut self,
state: State,
transition: Box<dyn Transition<Event, State, UserData> + Send + Sync>,
) {
self.transitions.insert(state, transition);
}
pub fn subscribe(&self) -> tokio::sync::broadcast::Receiver<State> {
self.broadcast.0.subscribe()
}
pub fn add_on_register_callback(
&mut self,
callback: FnOnEventRegister<Event, State, UserData>,
) {
self.on_event_register = Some(callback);
}
pub async fn process(&mut self) {
self.on_state_change();
while let Some(event) = self.event_receiver.recv().await {
self.register_event(event);
self.process_event(event).await;
self.broadcast.0.send(self.data.state).unwrap();
}
}
async fn process_event(&mut self, event: Event) {
if let Some(transition) = self.transitions.get_mut(&self.data.state) {
self.data.prev_state = Some(self.data.state);
self.data.state = transition.next(event.clone(), &mut self.data).await;
if self.data.prev_state.unwrap() != self.data.state {
self.on_state_change();
}
}
info!(
"[fsm] Processed event: {event:?}; {:?} => {:?}",
self.data.prev_state, self.data.state
);
}
fn register_event(&mut self, event: Event) {
self.data.events.insert(event, Instant::now());
if let Some(callback) = self.on_event_register {
(callback)(event, &mut self.data);
}
}
fn on_state_change(&mut self) {
if let Some(transition) = self.transitions.get_mut(&self.data.state) {
transition.enter(&mut self.data);
}
}
}
#[cfg(test)]
mod test {
use super::*;
use tokio::task::JoinHandle;
#[derive(Default, Debug, Eq, PartialEq, Copy, Clone, Hash)]
pub enum State {
#[default]
Idle,
State1,
State2,
}
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
pub enum Event {
Event1,
Event2,
Event3,
}
#[derive(Debug, Default)]
struct UserData {
event_counter: u64,
}
struct IdleState;
#[async_trait]
impl Transition<Event, State, UserData> for IdleState {
async fn next(&mut self, event: Event, data: &Data<Event, State, UserData>) -> State {
match event {
Event::Event1 => State::State1,
_ => data.state,
}
}
}
struct State1State;
#[async_trait]
impl Transition<Event, State, UserData> for State1State {
async fn next(&mut self, event: Event, data: &Data<Event, State, UserData>) -> State {
match event {
Event::Event2 => State::State2,
_ => data.state,
}
}
}
struct State2State;
#[async_trait]
impl Transition<Event, State, UserData> for State2State {
async fn next(&mut self, event: Event, data: &Data<Event, State, UserData>) -> State {
if data.user_data.event_counter > 5 {
return State::Idle;
}
match event {
Event::Event3 => State::Idle,
_ => data.state,
}
}
}
async fn create_stm() -> (
JoinHandle<()>,
tokio::sync::mpsc::Sender<Event>,
tokio::sync::broadcast::Receiver<State>,
) {
let (mut stm, event_sender) = StateMachine::<Event, State, UserData>::new(100);
stm.add_transition(State::Idle, Box::new(IdleState {}));
stm.add_transition(State::State1, Box::new(State1State {}));
stm.add_transition(State::State2, Box::new(State2State {}));
stm.add_on_register_callback(|_, data| {
data.user_data.event_counter = data.user_data.event_counter + 1;
});
let sub = stm.subscribe();
let task = tokio::spawn(async move {
let _ = stm.process().await;
});
(task, event_sender, sub)
}
#[tokio::test]
async fn given_idle_state_when_event1_occur_then_state_change_to_state1() {
let (task, sender, mut states) = create_stm().await;
let _ = sender.send(Event::Event1).await;
assert_eq!(states.recv().await.unwrap(), State::State1);
task.abort();
}
#[tokio::test]
async fn given_idle_state_when_event2_occur_then_state_remain_the_same() {
let (task, sender, mut states) = create_stm().await;
let _ = sender.send(Event::Event2).await;
assert_eq!(states.recv().await.unwrap(), State::Idle);
task.abort();
}
#[tokio::test]
async fn given_state1_when_event3_occur_then_state_return_to_idle() {
let (task, sender, mut states) = create_stm().await;
let _ = sender.send(Event::Event1).await;
assert_eq!(states.recv().await.unwrap(), State::State1);
let _ = sender.send(Event::Event2).await;
assert_eq!(states.recv().await.unwrap(), State::State2);
let _ = sender.send(Event::Event3).await;
assert_eq!(states.recv().await.unwrap(), State::Idle);
task.abort();
}
#[tokio::test]
async fn given_state2_when_events_counter_exceeded_then_state_return_to_idle() {
let (task, sender, mut states) = create_stm().await;
let _ = sender.send(Event::Event1).await;
assert_eq!(states.recv().await.unwrap(), State::State1);
let _ = sender.send(Event::Event2).await;
assert_eq!(states.recv().await.unwrap(), State::State2);
for _ in 0..3 {
let _ = sender.send(Event::Event1).await;
assert_eq!(states.recv().await.unwrap(), State::State2);
}
let _ = sender.send(Event::Event1).await;
assert_eq!(states.recv().await.unwrap(), State::Idle);
task.abort();
}
}