Skip to content

Commit 182cb55

Browse files
committed
optional choices
1 parent f76053c commit 182cb55

6 files changed

Lines changed: 234 additions & 120 deletions

File tree

dsm/__init__.py

Lines changed: 28 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,3 @@
1-
21
import collections
32
import observable
43

@@ -43,8 +42,9 @@ def has_state(self, state):
4342
def register(self, from_state, value, to_state):
4443
if from_state in self._states and value in self._states[from_state]:
4544
raise AlreadyRegistered(
46-
'Transition for `%s` is already registered for state `%s`' % (
47-
value, from_state))
45+
"Transition for `%s` is already registered for state `%s`"
46+
% (value, from_state)
47+
)
4848
self._states[from_state][value] = to_state
4949
self._allstates.update([from_state, to_state])
5050

@@ -60,16 +60,16 @@ def register_fallback(self, from_state, to_state):
6060

6161
if from_state in self._fallbacks:
6262
raise AlreadyRegistered(
63-
'Fallback transition for `%s` '
64-
'is already registered' % from_state)
63+
"Fallback transition for `%s` " "is already registered" % from_state
64+
)
6565

6666
self._fallbacks[from_state] = to_state
6767
self._allstates.update([from_state, to_state])
6868

6969
def can(self, value, current_state):
7070
return bool(
71-
self._states.get(current_state) and
72-
self._states[current_state].get(value))
71+
self._states.get(current_state) and self._states[current_state].get(value)
72+
)
7373

7474
def execute(self, value, current_state):
7575
try:
@@ -79,8 +79,9 @@ def execute(self, value, current_state):
7979
return self._fallbacks[current_state]
8080
except KeyError:
8181
raise UnknownTransition(
82-
'Can not find transition for `%s` in state `%s`' % (
83-
value, current_state))
82+
"Can not find transition for `%s` in state `%s`"
83+
% (value, current_state)
84+
)
8485

8586

8687
class MetaMachine(type):
@@ -93,36 +94,38 @@ def __new__(cls, name, bases, attrs):
9394
cls.add_exception_classes(new_class)
9495
return new_class
9596

96-
meta = attrs.pop('Meta', None)
97+
meta = attrs.pop("Meta", None)
9798

9899
class Options:
99100
def __init__(self, meta):
100101
self.transitions = Transitions(
101-
transitions=getattr(meta, 'transitions', None),
102-
fallbacks=getattr(meta, 'fallbacks', None))
103-
self.initial = getattr(meta, 'initial', None)
102+
transitions=getattr(meta, "transitions", None),
103+
fallbacks=getattr(meta, "fallbacks", None),
104+
)
105+
self.initial = getattr(meta, "initial", None)
104106

105107
new_class = super_new(cls, name, bases, {})
106108
cls.add_exception_classes(new_class)
107-
setattr(new_class, '_meta', Options(meta))
109+
setattr(new_class, "_meta", Options(meta))
108110

109111
return new_class
110112

111113
def add_exception_classes(new_class):
112-
setattr(new_class, 'FSMException', FSMException)
113-
setattr(new_class, 'UnknownTransition', UnknownTransition)
114+
setattr(new_class, "FSMException", FSMException)
115+
setattr(new_class, "UnknownTransition", UnknownTransition)
114116

115117

116118
class StateMachine(metaclass=MetaMachine):
117119
def __init__(self, initial=None, transitions=None):
118-
meta = getattr(self, '_meta', None)
120+
meta = getattr(self, "_meta", None)
119121
self._eventhandler = observable.Observable()
120-
self._transitions = transitions or getattr(
121-
meta, 'transitions', None) or Transitions()
122-
self._initial = initial or getattr(meta, 'initial', None)
122+
self._transitions = (
123+
transitions or getattr(meta, "transitions", None) or Transitions()
124+
)
125+
self._initial = initial or getattr(meta, "initial", None)
123126
self._state = None
124127
self._inputhandlers = collections.defaultdict(list)
125-
self._eventhandler.on('input', self._inputhandler)
128+
self._eventhandler.on("input", self._inputhandler)
126129
self.reset()
127130

128131
@property
@@ -133,11 +136,10 @@ def process(self, value):
133136
new_state = self._transitions.execute(value, self.state)
134137

135138
if not self.state == new_state:
136-
self._eventhandler.trigger(
137-
'change', state=new_state, previous=self.state)
139+
self._eventhandler.trigger("change", state=new_state, previous=self.state)
138140

139141
self._state = new_state
140-
self._eventhandler.trigger('input', state=new_state, value=value)
142+
self._eventhandler.trigger("input", state=new_state, value=value)
141143

142144
return self.state
143145

@@ -155,9 +157,8 @@ def reset(self):
155157

156158
old_state = self._state
157159
self._state = self._initial
158-
self._eventhandler.trigger(
159-
'change', state=self._state, previous=old_state)
160-
self._eventhandler.trigger('reset')
160+
self._eventhandler.trigger("change", state=self._state, previous=old_state)
161+
self._eventhandler.trigger("reset")
161162
return self.state
162163

163164
def when(self, state, func):

dsm/fields.py

Lines changed: 28 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,9 @@
1+
from django.core.checks import Error
12
from django.db import models
23

34
from . import StateMachine, Transitions
45

5-
__all__ = ['StateMachineField']
6+
__all__ = ["StateMachineField"]
67

78

89
class MachineState:
@@ -49,13 +50,11 @@ def __get__(self, instance, cls=None):
4950
return value
5051
if isinstance(value, MachineState):
5152
return value
52-
return MachineState(
53-
instance, self.field, self.field._create_fsm(value))
53+
return MachineState(instance, self.field, self.field._create_fsm(value))
5454

5555
def __set__(self, instance, value):
5656
if value is not None:
57-
value = MachineState(
58-
instance, self.field, self.field._create_fsm(value))
57+
value = MachineState(instance, self.field, self.field._create_fsm(value))
5958
instance.__dict__[self.field.name] = value
6059
"""
6160
current_value = instance.__dict__.get(self.field.name)
@@ -74,12 +73,18 @@ class StateMachineFieldMixin:
7473
descriptor_class = StateDescriptor
7574

7675
def __init__(self, transitions, *args, **kwargs):
77-
self.transitions = transitions
76+
if isinstance(transitions, Transitions):
77+
self.transitions = transitions
78+
else:
79+
self.transitions = Transitions(transitions)
80+
81+
if "choices" not in kwargs:
82+
kwargs["choices"] = [(s, s) for s in self.transitions._allstates]
83+
7884
super().__init__(*args, **kwargs)
7985

8086
def _create_fsm(self, initial):
81-
return StateMachine(
82-
initial=initial, transitions=Transitions(self.transitions))
87+
return StateMachine(initial=initial, transitions=self.transitions)
8388

8489
def get_prep_value(self, value):
8590
if value is None:
@@ -92,8 +97,22 @@ def contribute_to_class(self, cls, name, **kwargs):
9297

9398

9499
class StateMachineField(StateMachineFieldMixin, models.CharField):
100+
def _check_choices(self, **kwargs):
101+
all_states = self.transitions._allstates
102+
choice_states = {x[0] for x in self.choices}
103+
if not all_states.issubset(choice_states):
104+
return [
105+
Error(
106+
"Following states are not defined in choices: %s"
107+
% (", ".join(all_states - choice_states)),
108+
obj=self,
109+
id="dsm.E001",
110+
)
111+
]
112+
return []
113+
95114
def deconstruct(self):
96115
name, path, args, kwargs = super().deconstruct()
97-
path = 'dsm.fields.StateMachineField'
116+
path = "dsm.fields.StateMachineField"
98117
args.insert(0, self.transitions)
99118
return name, path, args, kwargs

requirements-dev.txt

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,3 +4,5 @@ ipdb
44
wheel
55
pytest
66
pytest-django
7+
isort
8+
black

tests/models.py

Lines changed: 30 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,14 @@
11
from django.db import models
2-
from django.utils.translation import gettext_lazy as _
32
from dsm.fields import StateMachineField
43

4+
55
class Order(models.Model):
66
class Status(models.TextChoices):
7-
NEW = 'new', _('New')
8-
PROCESSING = 'processing', _('Processing')
9-
SENDING = 'sending', _('Sending')
10-
FINISHED = 'finished', _('Finished')
11-
CANCELLED = 'cancelled', _('Cancelled')
7+
NEW = 'new', 'New'
8+
PROCESSING = 'processing', 'Processing'
9+
SENDING = 'sending', 'Sending'
10+
FINISHED = 'finished', 'Finished'
11+
CANCELLED = 'cancelled', 'Cancelled'
1212

1313
status = StateMachineField(
1414
transitions=(
@@ -25,3 +25,27 @@ class Status(models.TextChoices):
2525

2626
class Meta:
2727
app_label = 'tests'
28+
29+
30+
class OrderNoChoices(models.Model):
31+
class Status(models.TextChoices):
32+
NEW = 'new', 'New'
33+
PROCESSING = 'processing', 'Processing'
34+
SENDING = 'sending', 'Sending'
35+
FINISHED = 'finished', 'Finished'
36+
CANCELLED = 'cancelled', 'Cancelled'
37+
38+
status = StateMachineField(
39+
transitions=(
40+
(Status.NEW, ['confirm'], Status.PROCESSING),
41+
(Status.PROCESSING, ['cancel'], Status.CANCELLED),
42+
(Status.PROCESSING, ['send'], Status.SENDING),
43+
(Status.SENDING, ['deliver'], Status.FINISHED),
44+
),
45+
max_length=16,
46+
db_index=True,
47+
default=Status.NEW,
48+
)
49+
50+
class Meta:
51+
app_label = 'tests'

0 commit comments

Comments
 (0)