aboutsummaryrefslogtreecommitdiffstats
path: root/desmo_api/fsm.py
blob: dfbef20649379d5ba5df66350a39d545e4bb4912 (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
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
import asyncio
from collections.abc import Coroutine
import logging
from statemachine import StateMachine, State


from . import actions
from . import hcloud_dns
from . import db


logger = logging.getLogger(__name__)


class JailStateMachine(StateMachine):
    def __init__(
        self,
        dns_client: hcloud_dns.HCloudDNS,
        database: db.DB,
        name: str,
    ):
        self._dns_client = dns_client
        self._name = name
        self._db = database
        self._tasks = set()

        super().__init__()

    uninitialized = State(initial=True)
    jail_provisioning = State()
    dns_provisioning = State()
    jail_setup = State()
    jail_ready = State()
    jail_removal = State()
    dns_deprovisioning = State()
    terminated = State()

    initialize = uninitialized.to(jail_provisioning) | terminated.to(jail_provisioning)
    jail_provisioned = jail_provisioning.to(dns_provisioning)
    jail_provisioning_failed = jail_provisioning.to(jail_provisioning)
    dns_provisioned = dns_provisioning.to(jail_setup)
    dns_provisioning_failed = dns_provisioning.to(dns_provisioning)
    jail_setup_done = jail_setup.to(jail_ready)
    jail_setup_failed = jail_setup.to(jail_setup)
    remove_jail = jail_ready.to(jail_removal)
    jail_removal_failed = jail_removal.to(jail_removal)
    jail_removed = jail_removal.to(dns_deprovisioning)
    dns_deprovisioned = dns_deprovisioning.to(terminated)
    dns_deprovisioning_failed = dns_deprovisioning.to(dns_deprovisioning)

    def start_on_enter_task(self):
        state = self.current_state.id
        switch = {
            "jail_provisioning": self.on_enter_jail_provisioning,
            "dns_provisioning": self.on_enter_dns_provisioning,
            "jail_setup": self.on_enter_jail_setup,
            "jail_ready": self.on_enter_jail_ready,
            "jail_removal": self.on_enter_jail_removal,
            "dns_deprovisioning": self.on_enter_dns_deprovisioning,
        }
        func = switch.get(state)
        if func is not None:
            func()

    async def _store_state_and_run(self, awaitable: Coroutine):
        await self._db.set_jail_state(self._name, self.current_state.id)
        await awaitable

    def on_enter_jail_provisioning(self):
        task = asyncio.create_task(
            self._store_state_and_run(
                actions.start_jail_provisioning(self, self._db, self._name)
            )
        )
        self._tasks.add(task)

    def on_enter_dns_provisioning(self):
        task = asyncio.create_task(
            self._store_state_and_run(
                actions.start_dns_provisioning(
                    self, self._db, self._dns_client, self._name
                )
            )
        )
        self._tasks.add(task)

    def on_enter_jail_setup(self):
        task = asyncio.create_task(
            self._store_state_and_run(
                actions.start_jail_setup(self, self._db, self._name)
            )
        )
        self._tasks.add(task)

    def on_enter_jail_ready(self):
        task = asyncio.create_task(
            self._store_state_and_run(
                actions.start_jail_watch(self, self._db, self._name)
            )
        )
        self._tasks.add(task)

    def on_enter_jail_removal(self):
        task = asyncio.create_task(
            self._store_state_and_run(
                actions.start_jail_removal(self, self._db, self._name)
            )
        )
        self._tasks.add(task)

    def on_enter_dns_deprovisioning(self):
        task = asyncio.create_task(
            self._store_state_and_run(
                actions.start_dns_deprovisioning(
                    self, self._dns_client, self._db, self._name
                )
            )
        )
        self._tasks.add(task)

    def on_enter_terminated(self):
        logger.info("Jail %s terminated", self._name)
        for task in self._tasks:
            task.cancel()