summaryrefslogtreecommitdiff
path: root/tests/test_state.py
blob: 9fd223cacc4db8bd31cdfad75f1b1af5e046b55e (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
import unittest

from parley.state import InvalidStateError, State, StateMachine


class StateMachineTests(unittest.TestCase):
    def test_record_transcribe_complete(self) -> None:
        machine = StateMachine()
        self.assertEqual(machine.start_recording(), State.RECORDING)
        self.assertEqual(machine.begin_transcription(), State.TRANSCRIBING)
        self.assertEqual(machine.complete(), State.IDLE)

    def test_insertion_flow(self) -> None:
        machine = StateMachine()
        machine.start_recording()
        machine.begin_transcription()
        self.assertEqual(machine.begin_insertion(), State.INSERTING)
        self.assertEqual(machine.complete(), State.IDLE)

    def test_rejects_second_recording(self) -> None:
        machine = StateMachine()
        machine.start_recording()
        with self.assertRaisesRegex(InvalidStateError, "current state is recording"):
            machine.start_recording()

    def test_cancel_recording_and_transcription(self) -> None:
        machine = StateMachine()
        machine.start_recording()
        self.assertEqual(machine.cancel(), State.IDLE)
        machine.start_recording()
        machine.begin_transcription()
        self.assertEqual(machine.cancel(), State.IDLE)

    def test_error_is_recoverable(self) -> None:
        machine = StateMachine()
        machine.fail()
        self.assertEqual(machine.start_recording(), State.RECORDING)


if __name__ == "__main__":
    unittest.main()