92% test coverage
This commit is contained in:
@@ -0,0 +1,3 @@
|
|||||||
|
[run]
|
||||||
|
omit =
|
||||||
|
tests/*
|
||||||
+36
-36
@@ -20,16 +20,12 @@ TransitionInput = TypeVar("TransitionInput", bound=Union[InputProtocol, Transiti
|
|||||||
|
|
||||||
|
|
||||||
class State:
|
class State:
|
||||||
def __init__(self, name: str, data: object) -> None:
|
def __init__(self, name: str) -> None:
|
||||||
self.__name = name
|
self.__name = name
|
||||||
self.__data = data
|
|
||||||
|
|
||||||
def get_name(self) -> str:
|
def get_name(self) -> str:
|
||||||
return self.__name
|
return self.__name
|
||||||
|
|
||||||
def get_data(self) -> object:
|
|
||||||
return self.__data
|
|
||||||
|
|
||||||
def __str__(self) -> str:
|
def __str__(self) -> str:
|
||||||
return f"[State {self.__name}]"
|
return f"[State {self.__name}]"
|
||||||
|
|
||||||
@@ -37,9 +33,9 @@ class State:
|
|||||||
return hash(self.__name)
|
return hash(self.__name)
|
||||||
|
|
||||||
|
|
||||||
class Transition(Generic[Input]):
|
class Transition(Generic[TransitionInput]):
|
||||||
def __init__(self, transition_input: TransitionInput, origin_name: str, destination_name: str,
|
def __init__(self, transition_input: TransitionInput, origin_name: str, destination_name: str,
|
||||||
output_function: Callable[[Self, Optional[object]], None]) -> None:
|
output_function: Callable[[Optional[object]], None]) -> None:
|
||||||
self.__transition_input = transition_input
|
self.__transition_input = transition_input
|
||||||
self.__origin_name = origin_name
|
self.__origin_name = origin_name
|
||||||
self.__destination_name = destination_name
|
self.__destination_name = destination_name
|
||||||
@@ -47,7 +43,7 @@ class Transition(Generic[Input]):
|
|||||||
self.__output_function = output_function
|
self.__output_function = output_function
|
||||||
|
|
||||||
def execute_output(self, context):
|
def execute_output(self, context):
|
||||||
self.__output_function(context, self)
|
self.__output_function(context)
|
||||||
|
|
||||||
def get_destination_name_hash(self) -> int:
|
def get_destination_name_hash(self) -> int:
|
||||||
return self.__destination_name_hash
|
return self.__destination_name_hash
|
||||||
@@ -61,17 +57,20 @@ class Transition(Generic[Input]):
|
|||||||
|
|
||||||
|
|
||||||
class Machine(Generic[Input]):
|
class Machine(Generic[Input]):
|
||||||
def __init__(self) -> None:
|
def __init__(self, initial_context=None) -> None:
|
||||||
self.__start_state: Optional[int] = None
|
if initial_context is None:
|
||||||
|
initial_context = {}
|
||||||
|
self.__start_state_hash: Optional[int] = None
|
||||||
self.__end_states: List[int] = []
|
self.__end_states: List[int] = []
|
||||||
self.__current_state: Optional[int] = None
|
self.__current_state_hash: Optional[int] = None
|
||||||
self.__states: dict[int, State] = {}
|
self.__states: dict[int, State] = {}
|
||||||
# dict[hash_origin, dict[hash_transition_input, Transition]]
|
# dict[hash_origin, dict[hash_transition_input, Transition]]
|
||||||
self.__transitions: dict[int, dict[int, Transition]] = {}
|
self.__transitions: dict[int, dict[int, Transition]] = {}
|
||||||
self.__context = {}
|
self.__initial_context = initial_context
|
||||||
|
self.__context = initial_context
|
||||||
|
|
||||||
def add_transition(self, transition_input: TransitionInput, origin_name: str, destination_name: str,
|
def add_transition(self, transition_input: TransitionInput, origin_name: str, destination_name: str,
|
||||||
output_function: Callable[[Self, object], None]) -> None:
|
output_function: Callable[[Optional[object]], None]) -> None:
|
||||||
ho = hash(origin_name)
|
ho = hash(origin_name)
|
||||||
if ho not in self.__states.keys():
|
if ho not in self.__states.keys():
|
||||||
raise ValueError(f"Origin state {origin_name} does not exist.")
|
raise ValueError(f"Origin state {origin_name} does not exist.")
|
||||||
@@ -90,14 +89,13 @@ class Machine(Generic[Input]):
|
|||||||
elif hi in self.__transitions[ho]:
|
elif hi in self.__transitions[ho]:
|
||||||
raise ValueError(f"Transition already exists for origin state {origin_name} and transition input {transition_input}.")
|
raise ValueError(f"Transition already exists for origin state {origin_name} and transition input {transition_input}.")
|
||||||
|
|
||||||
def add_state(self, state_name: str, state_data: object, is_start_state: bool = False,
|
def add_state(self, state_name: str, is_start_state: bool = False, is_end_state: bool = False) -> None:
|
||||||
is_end_state: bool = False) -> None:
|
|
||||||
if state_name.startswith("#"):
|
if state_name.startswith("#"):
|
||||||
raise ValueError("State names cannot start with # (reserved for special states)")
|
raise ValueError("State names cannot start with # (reserved for special states)")
|
||||||
h: int = hash(state_name)
|
h: int = hash(state_name)
|
||||||
if is_start_state:
|
if is_start_state:
|
||||||
if self.__start_state is None:
|
if self.__start_state_hash is None:
|
||||||
self.__start_state = h
|
self.__start_state_hash = h
|
||||||
else:
|
else:
|
||||||
raise ValueError("Adding a start state when one is already set.")
|
raise ValueError("Adding a start state when one is already set.")
|
||||||
if is_end_state:
|
if is_end_state:
|
||||||
@@ -105,7 +103,7 @@ class Machine(Generic[Input]):
|
|||||||
self.__end_states.append(h)
|
self.__end_states.append(h)
|
||||||
if h in self.__states.keys():
|
if h in self.__states.keys():
|
||||||
raise ValueError(f"State {state_name} already exists.")
|
raise ValueError(f"State {state_name} already exists.")
|
||||||
state = State(state_name, state_data)
|
state = State(state_name)
|
||||||
self.__states[h] = state
|
self.__states[h] = state
|
||||||
|
|
||||||
def get_state(self, name: str) -> Optional[State]:
|
def get_state(self, name: str) -> Optional[State]:
|
||||||
@@ -113,28 +111,18 @@ class Machine(Generic[Input]):
|
|||||||
# KeyError if the state is not found
|
# KeyError if the state is not found
|
||||||
return self.__states[h]
|
return self.__states[h]
|
||||||
|
|
||||||
def reset(self) -> None:
|
|
||||||
if self.__start_state is None:
|
|
||||||
raise ValueError("Resetting a machine without starting state.")
|
|
||||||
self.__current_state = self.__start_state
|
|
||||||
|
|
||||||
def get_current_state(self) -> Optional[State]:
|
|
||||||
if self.__current_state is None:
|
|
||||||
return None # The machine has not started yet
|
|
||||||
return self.__states[self.__current_state]
|
|
||||||
|
|
||||||
def process_input(self, i: Input) -> None:
|
def process_input(self, i: Input) -> None:
|
||||||
if self.__start_state is None:
|
if self.__start_state_hash is None:
|
||||||
raise ValueError("Processing input without a start state.")
|
raise ValueError("Processing input without a start state.")
|
||||||
if self.__current_state is None:
|
if self.__current_state_hash is None:
|
||||||
self.__current_state = self.__start_state
|
self.__current_state_hash = self.__start_state_hash
|
||||||
hs = hash(self.__current_state)
|
|
||||||
# dict[hash_origin, dict[hash_transition_input, Transition]]
|
# dict[hash_origin, dict[hash_transition_input, Transition]]
|
||||||
if hs not in self.__transitions.keys():
|
if self.__current_state_hash not in self.__transitions.keys():
|
||||||
raise ValueError("Processing input for a state with no transitions.")
|
raise ValueError("Processing input for a state with no transitions.")
|
||||||
transitions = self.__transitions[hs]
|
# noinspection PyTypeChecker
|
||||||
|
transitions = self.__transitions[self.__current_state_hash]
|
||||||
hi = hash(i)
|
hi = hash(i)
|
||||||
transition = None
|
transition: Optional[Transition] = None
|
||||||
if hi in transitions.keys():
|
if hi in transitions.keys():
|
||||||
# Transition found: execute and move to the destination state
|
# Transition found: execute and move to the destination state
|
||||||
transition = transitions[hi]
|
transition = transitions[hi]
|
||||||
@@ -145,8 +133,20 @@ class Machine(Generic[Input]):
|
|||||||
transition = transitions[hc]
|
transition = transitions[hc]
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Found invalid transition while processing input {i} [{transition}].]")
|
raise ValueError(f"Found invalid transition while processing input {i} [{transition}].]")
|
||||||
|
assert transition is not None
|
||||||
transition.execute_output(self.__context)
|
transition.execute_output(self.__context)
|
||||||
self.__current_state = transition.get_destination_name_hash()
|
self.__current_state_hash = transition.get_destination_name_hash()
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
if self.__start_state_hash is None:
|
||||||
|
raise ValueError("Resetting a machine without starting state.")
|
||||||
|
self.__current_state_hash = self.__start_state_hash
|
||||||
|
self.__context = self.__initial_context
|
||||||
|
|
||||||
|
def get_current_state(self) -> Optional[State]:
|
||||||
|
if self.__current_state_hash is None:
|
||||||
|
return None # The machine has not started yet
|
||||||
|
return self.__states[self.__current_state_hash]
|
||||||
|
|
||||||
def __str__(self) -> str:
|
def __str__(self) -> str:
|
||||||
r = [f"=MACHINE=", "\tSTATES"]
|
r = [f"=MACHINE=", "\tSTATES"]
|
||||||
|
|||||||
@@ -1,21 +0,0 @@
|
|||||||
import unittest
|
|
||||||
|
|
||||||
from src.snanosm.mealy import Machine
|
|
||||||
|
|
||||||
|
|
||||||
class TestMachine(unittest.TestCase):
|
|
||||||
|
|
||||||
def test_init(self):
|
|
||||||
m = Machine()
|
|
||||||
self.assertIsNotNone(m)
|
|
||||||
|
|
||||||
def test_add_state_two_start_states(self):
|
|
||||||
m = Machine()
|
|
||||||
m.add_state("A", {}, True, False)
|
|
||||||
self.assertRaises(ValueError, m.add_state, "B", {}, True, False)
|
|
||||||
|
|
||||||
def test_add_two_end_states(self):
|
|
||||||
m = Machine()
|
|
||||||
m.add_state("A", {}, True, False)
|
|
||||||
m.add_state("B", {}, False, True)
|
|
||||||
m.add_state("C", {}, False, True)
|
|
||||||
@@ -0,0 +1,236 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from snanosm.mealy import Machine, TransitionInputEnum, State, Transition
|
||||||
|
|
||||||
|
|
||||||
|
# Any hashable and equatable object should suffice
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class TestTransitionObject:
|
||||||
|
attribute: str
|
||||||
|
|
||||||
|
|
||||||
|
class TestMachine(unittest.TestCase):
|
||||||
|
|
||||||
|
def test_init(self):
|
||||||
|
m = Machine()
|
||||||
|
self.assertIsNotNone(m)
|
||||||
|
|
||||||
|
# ADD STATE
|
||||||
|
|
||||||
|
def test_add_state_two_start_states(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.add_state("B", True, False)
|
||||||
|
|
||||||
|
def test_add_two_end_states(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_state("C", False, True)
|
||||||
|
|
||||||
|
def test_add_state_with_reserved_name(self):
|
||||||
|
m = Machine()
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.add_state("#INVALID_STATE_NAME", True, False)
|
||||||
|
|
||||||
|
def test_add_state_with_duplicate_state_name(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.add_state("A", False, True)
|
||||||
|
|
||||||
|
# GET STATE
|
||||||
|
|
||||||
|
def test_get_states(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
state = m.get_state("A")
|
||||||
|
self.assertIsNotNone(state)
|
||||||
|
assert state is not None
|
||||||
|
self.assertEqual(state.get_name(), "A")
|
||||||
|
|
||||||
|
# ADD TRANSITION
|
||||||
|
|
||||||
|
def test_add_transition(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print(context))
|
||||||
|
|
||||||
|
def test_add_transition_non_existent_origin_state(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.add_transition("X", "C", "B", lambda context: print(context))
|
||||||
|
|
||||||
|
def test_add_transition_non_existent_destination_state(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.add_transition("X", "A", "C", lambda context: print(context))
|
||||||
|
|
||||||
|
def test_add_special_transition(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition(TransitionInputEnum.MATCH_REST, "A", "B", lambda context: print(context))
|
||||||
|
|
||||||
|
def test_add_several_transitions_with_same_origin(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("1"))
|
||||||
|
m.add_transition("Y", "A", "B", lambda context: print("2"))
|
||||||
|
|
||||||
|
def test_add_duplicate_transition(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("1"))
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("2"))
|
||||||
|
|
||||||
|
# PROCESS INPUT
|
||||||
|
|
||||||
|
def test_process_input(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, False)
|
||||||
|
m.add_state("C", False, True)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("TEST_PROCESS_INPUT: Transition 1"))
|
||||||
|
m.add_transition("X", "B", "C", lambda context: print("TEST_PROCESS_INPUT: Transition 2"))
|
||||||
|
m.process_input("X")
|
||||||
|
m.process_input("X")
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "C")
|
||||||
|
|
||||||
|
def test_process_input_no_start_state(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", False, False)
|
||||||
|
m.add_state("B", False, False)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("TEST_PROCESS_INPUT_NO_START_DATE: Transition"))
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.process_input("X")
|
||||||
|
|
||||||
|
def test_process_input_state_without_transitions(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("TEST_PROCESS_INPUT_WITHOUT_TRANSITIONS: Transition"))
|
||||||
|
m.process_input("X")
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "B")
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.process_input("X")
|
||||||
|
|
||||||
|
def test_process_input_catch_all_transition(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition(TransitionInputEnum.MATCH_REST, "A", "B", lambda context: print("A -> B"))
|
||||||
|
m.add_transition(TransitionInputEnum.MATCH_REST, "B", "A", lambda context: print("B -> A"))
|
||||||
|
m.process_input("X")
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "B")
|
||||||
|
m.process_input(1)
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "A")
|
||||||
|
m.process_input(TestTransitionObject("X"))
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "B")
|
||||||
|
|
||||||
|
def test_process_input_with_objects(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, False)
|
||||||
|
m.add_state("C", False, True)
|
||||||
|
m.add_transition(TestTransitionObject("TEST_1"), "A", "B", lambda context: print("A -> B"))
|
||||||
|
m.add_transition(TestTransitionObject("TEST_2"), "B", "C", lambda context: print("B -> C"))
|
||||||
|
m.process_input(TestTransitionObject("TEST_1"))
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "B")
|
||||||
|
m.process_input(TestTransitionObject("TEST_2"))
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "C")
|
||||||
|
|
||||||
|
def test_process_input_invalid_transition(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("A -> B"))
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.process_input("Y")
|
||||||
|
|
||||||
|
# GET CURRENT STATE
|
||||||
|
def test_get_current_state_machine_not_started(self):
|
||||||
|
m = Machine()
|
||||||
|
state = m.get_current_state()
|
||||||
|
self.assertIsNone(state)
|
||||||
|
|
||||||
|
def test_get_current_state_machine_started(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("A -> B"))
|
||||||
|
m.process_input("X")
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "B")
|
||||||
|
|
||||||
|
# RESET
|
||||||
|
|
||||||
|
def test_reset(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", True, False)
|
||||||
|
m.add_state("B", False, True)
|
||||||
|
m.add_transition("X", "A", "B", lambda context: print("A -> B"))
|
||||||
|
m.process_input("X")
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "B")
|
||||||
|
m.reset()
|
||||||
|
current_state = m.get_current_state()
|
||||||
|
self.assertIsNotNone(current_state)
|
||||||
|
assert current_state is not None
|
||||||
|
self.assertEqual(current_state.get_name(), "A")
|
||||||
|
|
||||||
|
def test_reset_no_start_state(self):
|
||||||
|
m = Machine()
|
||||||
|
m.add_state("A", False, False)
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
m.reset()
|
||||||
|
|
||||||
|
# STATE
|
||||||
|
|
||||||
|
def test_state_str(self):
|
||||||
|
s = State("A")
|
||||||
|
t = str(s)
|
||||||
|
self.assertEqual(t, "[State A]")
|
||||||
|
|
||||||
|
# TRANSITION
|
||||||
|
|
||||||
|
def test_transition_str(self):
|
||||||
|
t = Transition("TEST_1", "A", "B", lambda context: print("A -> B"))
|
||||||
|
self.assertEqual(str(t), "[Transition (A, B, TEST_1)]")
|
||||||
|
|
||||||
|
|
||||||
Reference in New Issue
Block a user