diff options
Diffstat (limited to 'desmo_api')
| -rw-r--r-- | desmo_api/__init__.py | 1 | ||||
| -rw-r--r-- | desmo_api/__main__.py | 0 | ||||
| -rw-r--r-- | desmo_api/actions.py | 205 | ||||
| -rw-r--r-- | desmo_api/api.py | 112 | ||||
| -rw-r--r-- | desmo_api/db.py | 138 | ||||
| -rw-r--r-- | desmo_api/fsm.py | 124 | ||||
| -rw-r--r-- | desmo_api/hcloud_dns.py | 146 | ||||
| -rw-r--r-- | desmo_api/models.py | 32 |
8 files changed, 758 insertions, 0 deletions
diff --git a/desmo_api/__init__.py b/desmo_api/__init__.py new file mode 100644 index 0000000..489dadd --- /dev/null +++ b/desmo_api/__init__.py @@ -0,0 +1 @@ +from .api import app diff --git a/desmo_api/__main__.py b/desmo_api/__main__.py new file mode 100644 index 0000000..e69de29 --- /dev/null +++ b/desmo_api/__main__.py diff --git a/desmo_api/actions.py b/desmo_api/actions.py new file mode 100644 index 0000000..767acbc --- /dev/null +++ b/desmo_api/actions.py @@ -0,0 +1,205 @@ +from __future__ import annotations + +import asyncio +import logging +import os + +import ansible_runner +from typing import TYPE_CHECKING, Optional + +if TYPE_CHECKING: + from .fsm import JailStateMachine + +from . import hcloud_dns, db, models + +logger = logging.getLogger(__name__) + + +def get_inventory(vars: Optional[dict] = None) -> dict: + return { + "all": { + "vars": { + "ansible_user": "automation", + "ansible_python_interpreter": "/usr/local/bin/python", + "jails_path": "/usr/local/jails", + "media_path": "/usr/local/jails/media", + "containers_path": "/usr/local/jails/containers", + "ansible_ssh_common_args": "-o StrictHostKeyChecking=no", + **(vars or {}), + }, + "children": { + "bsd_servers": { + "hosts": { + "bsd-1.hki-rok.atk.works": {}, + "bsd-2.hki-rok.atk.works": {}, + "bsd-3.hki-rok.atk.works": {}, + }, + } + }, + } + } + + +def run_ansible_playbook(name, playbook: str, inventory: dict): + data_dir = f"/tmp/ansible-isolation/{name}" + os.makedirs(data_dir, exist_ok=True) + with open(os.environ["SSH_KEY_FILE"], "r") as f: + ssh_key = f.read() + return ansible_runner.run_async( + private_data_dir=data_dir, + project_dir=f"{os.getcwd()}/ansible/project", + playbook=playbook, + inventory=inventory, + ssh_key=ssh_key, + ) + + +async def jail_provisioning(jail_info: models.JailInfo): + vars = { + "jail_host": jail_info.host, + "jail_name": jail_info.name, + "jail_ipv6": jail_info.ip, + } + inventory = get_inventory(vars) + _thread, runner = run_ansible_playbook( + jail_info.name, "create_jail.yaml", inventory + ) + while runner.rc is None: + await asyncio.sleep(1) + if runner.rc != 0: + raise Exception(f"Ansible failed with {runner.rc}") + + +async def start_jail_provisioning(fsm: JailStateMachine, database: db.DB, name: str): + logger.info("Starting provisioning for jail %s", name) + try: + jail_info = await database.get_jail(name) + await jail_provisioning(jail_info=jail_info) + fsm.jail_provisioned() + except Exception as e: + logger.error("Jail provisioning failed for %s", name, exc_info=e) + fsm.jail_provisioning_failed() + + +async def dns_provisioning( + dns_client: hcloud_dns.HCloudDNS, jail_info: models.JailInfo +): + zone_id = os.environ["HCLOUD_DNS_ZONE_ID"] + name = jail_info.name + ipv6 = jail_info.ip + records = await dns_client.get_records_by_name(zone_id, name) + for record in records: + logger.info("Deleting existing record %s for jail %s", record.id, name) + await dns_client.delete_record(record.id) + + logger.info("Creating AAAA record for jail %s", name) + await dns_client.create_record( + zone_id=zone_id, + name=f"{name}.jail", + record_type="AAAA", + value=ipv6, + ) + + +async def start_dns_provisioning( + fsm: JailStateMachine, + database: db.DB, + dns_client: hcloud_dns.HCloudDNS, + name: str, +): + logger.info("Starting DNS provisioning for jail %s", name) + try: + jail_info = await database.get_jail(name) + await dns_provisioning(dns_client, jail_info) + fsm.dns_provisioned() + except Exception as e: + logger.error("DNS provisioning failed for server %s", name, exc_info=e) + fsm.dns_provisioning_failed() + + +async def jail_setup( + jail_info: models.JailInfo, packages: list[str], commands: list[str] +): + vars = { + "jail_host": jail_info.host, + "jail_name": jail_info.name, + "jail_packages": packages, + "jail_commands": commands, + } + inventory = get_inventory(vars) + _thread, runner = run_ansible_playbook(jail_info.name, "setup_jail.yaml", inventory) + while runner.rc is None: + await asyncio.sleep(1) + if runner.rc != 0: + raise Exception(f"Ansible failed with {runner.rc}") + + +async def start_jail_setup(fsm: JailStateMachine, database: db.DB, name: str): + logger.info("Starting jail setup for jail %s", name) + try: + jail_info = await database.get_jail(name) + packages = await database.get_jail_packages(name) + commands = await database.get_jail_commands(name) + await jail_setup(jail_info=jail_info, packages=packages, commands=commands) + fsm.jail_setup_done() + except Exception as e: + logger.error("Jail setup failed for jail %s", name, exc_info=e) + fsm.jail_setup_failed() + + +async def start_jail_watch(fsm: JailStateMachine, database: db.DB, name: str): + logger.info("Starting healtcheck for jail %s", name) + try: + logger.info("not implemented lol") + except Exception as e: + logger.error("Server healtcheck failed for jail %s", name, exc_info=e) + + +async def jail_removal(jail_info: models.JailInfo): + vars = { + "jail_host": jail_info.host, + "jail_name": jail_info.name, + } + inventory = get_inventory(vars) + _thread, runner = run_ansible_playbook( + jail_info.name, "delete_jail.yaml", inventory + ) + while runner.rc is None: + await asyncio.sleep(1) + if runner.rc != 0: + raise Exception(f"Ansible failed with {runner.rc}") + + +async def start_jail_removal(fsm: JailStateMachine, database: db.DB, name: str): + logger.info("Starting removal of jail %s", name) + try: + jail_info = await database.get_jail(name) + await jail_removal(jail_info) + fsm.jail_removed() + except Exception as e: + logger.error("Stop server failed for server %s, ignoring", name, exc_info=e) + fsm.jail_removal_failed() + + +async def dns_deprovisioning(dns_client: hcloud_dns.HCloudDNS, name: str): + zone_id = os.environ["HCLOUD_DNS_ZONE_ID"] + + records = await dns_client.get_records_by_name(zone_id, f"{name}.jail") + for record in records: + logger.info("Deleting record %s for jail %s", record.id, name) + await dns_client.delete_record(record.id) + + +async def start_dns_deprovisioning( + fsm: JailStateMachine, dns_client: hcloud_dns.HCloudDNS, database: db.DB, name: str +): + logger.info("Starting DNS deprovisioning for jail %s", name) + try: + await dns_deprovisioning(dns_client, name) + await database.set_jail_state( + name, "terminated" + ) # Have to do this here because the tasks will be cancelled + fsm.dns_deprovisioned() + except Exception as e: + logger.error("DNS deprovisioning failed for jail %s", name, exc_info=e) + fsm.dns_deprovisioning_failed() diff --git a/desmo_api/api.py b/desmo_api/api.py new file mode 100644 index 0000000..750904d --- /dev/null +++ b/desmo_api/api.py @@ -0,0 +1,112 @@ +from fastapi import FastAPI +from typing import Dict +from .fsm import JailStateMachine +import logging +import sys +import asyncio +import statemachine.exceptions +from contextlib import asynccontextmanager +from .hcloud_dns import HCloudDNS +import os +from . import db, models +import random +import secrets + +logging.basicConfig(stream=sys.stderr, level=logging.INFO) +logger = logging.getLogger("api") + +STATE_MACHINES: Dict[str, JailStateMachine] = {} + +DNS_CLIENT = HCloudDNS(os.environ["HCLOUD_DNS_KEY"]) +database = db.DB(os.environ["DATABASE_DSN"]) + + +@asynccontextmanager +async def lifespan(app: FastAPI): + logger.info("Running migrations") + await database.migrate() + logger.info("Loading servers from database") + jails = await database.get_jails() + logger.info("Loaded %s jails from database", len(jails)) + logger.info(jails) + for jail in jails: + logger.info("Loading server %s with state %s", jail.name, jail.state) + _fsm = JailStateMachine(DNS_CLIENT, database, jail.name) + _fsm.current_state_value = jail.state + _fsm.start_on_enter_task() + STATE_MACHINES[jail.name] = _fsm + yield + logger.info("Closing asyncio clients") + await DNS_CLIENT.close() + await database.close() + logger.info("Closed asyncio clients") + + +app = FastAPI(lifespan=lifespan) + + +@app.get("/") +async def read_root(): + return {"Hello": "World"} + + +@app.post( + "/jails", + status_code=201, +) +async def create_jail(req: models.CreateJailRequest) -> models.FullJailInfo: + first_part = secrets.token_hex(2) + second_part = secrets.token_hex(2) + name = f"{req.name}-{first_part}-{second_part}" + ip = os.environ["NETWORK_PREFIX"] + "::" + first_part + ":" + second_part + hosts = os.environ["RUNNER_HOSTS"].split(",") + host = random.choice(hosts) + state = "uninitialized" + + await database.insert_jail(name, host, ip, state) + for package in req.packages: + await database.insert_jail_package(name, package) + + for i, command in enumerate(req.commands): + await database.insert_jail_command(name, command, i) + + STATE_MACHINES[name] = JailStateMachine(DNS_CLIENT, database, name) + STATE_MACHINES[name].initialize() + await asyncio.sleep(1) + return models.FullJailInfo( + name=name, + state=state, + ip=ip, + host=host, + packages=req.packages, + commands=req.commands, + ) + + +@app.get("/jails/{name}") +async def get_server(name: str) -> models.FullJailInfo | Dict[str, str]: + if name not in STATE_MACHINES: + return {"error": "server does not exist"} + jail = await database.get_jail(name) + packages = await database.get_jail_packages(name) + commands = await database.get_jail_commands(name) + return models.FullJailInfo( + name=jail.name, + state=jail.state, + ip=jail.ip, + host=jail.host, + packages=packages, + commands=commands, + ) + + +@app.delete("/jails/{name}") +async def delete_server(name: str) -> Dict[str, str]: + if name not in STATE_MACHINES: + return {"error": "jail does not exist"} + try: + STATE_MACHINES[name].remove_jail() + except statemachine.exceptions.TransitionNotAllowed: + return {"error": "jail is not ready"} + await asyncio.sleep(1) + return {"status": "ok"} diff --git a/desmo_api/db.py b/desmo_api/db.py new file mode 100644 index 0000000..f326104 --- /dev/null +++ b/desmo_api/db.py @@ -0,0 +1,138 @@ +from typing import List, Optional + +import logging +import asyncpg + +from . import models + +logger = logging.getLogger(__name__) + +MIGRATIONS = [ + [ + """ + CREATE TABLE meta ( + version integer PRIMARY KEY + ); + """, + "INSERT INTO meta (version) VALUES (0)", + ], + [ + """ + CREATE TABLE jail ( + name text PRIMARY KEY, + host text NOT NULL, + ip text NOT NULL UNIQUE, + state text NOT NULL + ); + """, + """ + CREATE TABLE jail_package ( + jail_name text NOT NULL REFERENCES jail(name) ON DELETE CASCADE, + name text NOT NULL + ); + """, + """ + CREATE TABLE jail_command ( + jail_name text NOT NULL REFERENCES jail(name) ON DELETE CASCADE, + command text NOT NULL, + order_no integer NOT NULL + ); + """, + ], +] + + +class DB: + def __init__(self, dsn: str): + self.dsn = dsn + self._conn: Optional[asyncpg.Connection] = None + + async def close(self) -> None: + if self._conn is not None: + await self._conn.close() + + async def _get_conn(self) -> asyncpg.Connection: + if self._conn is None: + self._conn = await asyncpg.connect(self.dsn) + return self._conn + + async def migrate(self): + conn = await self._get_conn() + # get version from meta table + try: + version = await conn.fetchval("SELECT MAX(version) FROM meta") + except asyncpg.UndefinedTableError: + version = -1 + + assert type(version) is int, "version must be an integer" + + logger.info("Current database version: %s", version) + + for i, migration in enumerate(MIGRATIONS[version + 1 :]): + async with conn.transaction(): + for query in migration: + logger.info("Running migration: %s", query) + await conn.execute(query) + await conn.execute( + "INSERT INTO meta (version) VALUES ($1)", version + i + 1 + ) + + async def get_jails(self) -> List[models.JailInfo]: + conn = await self._get_conn() + rows = await conn.fetch("SELECT name, state, ip, host FROM jail;") + return [models.JailInfo(**row) for row in rows] + + async def get_jail(self, name: str) -> models.JailInfo: + conn = await self._get_conn() + row = await conn.fetchrow( + "SELECT name, state, ip, host FROM jail WHERE name = $1;", name + ) + return models.JailInfo(**row) + + async def get_jail_packages(self, name: str) -> List[str]: + conn = await self._get_conn() + rows = await conn.fetch( + "SELECT name FROM jail_package WHERE jail_name = $1;", name + ) + return [row["name"] for row in rows] + + async def get_jail_commands(self, name: str) -> List[str]: + conn = await self._get_conn() + rows = await conn.fetch( + "SELECT command FROM jail_command WHERE jail_name = $1 ORDER BY order_no;", + name, + ) + return [row["command"] for row in rows] + + async def delete_jail(self, name: str) -> None: + conn = await self._get_conn() + await conn.execute("DELETE FROM jail WHERE name = $1;", name) + + async def set_jail_state(self, name: str, state: str) -> None: + conn = await self._get_conn() + await conn.execute("UPDATE jail SET state = $1 WHERE name = $2;", state, name) + + async def insert_jail(self, name, host, ip, state) -> None: + conn = await self._get_conn() + await conn.execute( + "INSERT INTO jail (name, host, ip, state) VALUES ($1, $2, $3, $4);", + name, + host, + ip, + state, + ) + + async def insert_jail_package(self, name: str, package: str) -> None: + conn = await self._get_conn() + await conn.execute( + "INSERT INTO jail_package (jail_name, name) VALUES ($1, $2);", name, package + ) + + async def insert_jail_command(self, name: str, command: str, order: int) -> None: + conn = await self._get_conn() + await conn.execute( + "INSERT INTO jail_command (jail_name, command, order_no) VALUES ($1, $2, $3);", + name, + command, + order, + ) diff --git a/desmo_api/fsm.py b/desmo_api/fsm.py new file mode 100644 index 0000000..dfbef20 --- /dev/null +++ b/desmo_api/fsm.py @@ -0,0 +1,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() diff --git a/desmo_api/hcloud_dns.py b/desmo_api/hcloud_dns.py new file mode 100644 index 0000000..e8b6a3a --- /dev/null +++ b/desmo_api/hcloud_dns.py @@ -0,0 +1,146 @@ +import asyncio +import aiohttp +from typing import Dict, List, Optional + +from pydantic import BaseModel + + +class TxtVerification(BaseModel): + name: str + token: str + + +class ZoneResponse(BaseModel): + id: str + created: str + modified: str + legacy_dns_host: str + legacy_ns: List[str] + name: str + ns: List[str] + owner: str + paused: bool + permission: str + project: str + registrar: str + status: str + ttl: int + verified: str + records_count: int + is_secondary_dns: bool + txt_verification: TxtVerification + + +class RecordResponse(BaseModel): + type: str + id: str + created: str + modified: str + zone_id: str + name: str + value: str + ttl: int | None = None + + +class HCloudDNS: + def __init__(self, token: str, api_domain="dns.hetzner.com"): + self._token = token + self.api_domain = api_domain + self._session: Optional[aiohttp.ClientSession] = None + + def _get_session(self): + if self._session is None: + self._session = aiohttp.ClientSession( + headers={ + "Auth-API-Token": self._token, + "Content-Type": "application/json; charset=utf-8", + } + ) + return self._session + + async def close(self): + if self._session is not None: + await self._session.close() + await asyncio.sleep(0.250) + + async def get_all_zones(self, name: Optional[str] = None) -> List[ZoneResponse]: + session = self._get_session() + params = {} + if name is not None: + params["name"] = name + async with session.get( + f"https://{ self.api_domain }/api/v1/zones", params=params + ) as resp: + resp.raise_for_status() + return [ZoneResponse(**zone) for zone in (await resp.json())["zones"]] + + async def get_zone(self, zone_id: str) -> ZoneResponse: + session = self._get_session() + async with session.get( + f"https://{ self.api_domain }/api/v1/zones/{zone_id}" + ) as resp: + resp.raise_for_status() + return ZoneResponse(**(await resp.json())["zone"]) + + async def get_zone_by_name(self, name: str) -> ZoneResponse: + zones = await self.get_all_zones(name) + if len(zones) == 0: + raise ValueError(f"Zone '{name}' not found") + elif len(zones) > 1: + raise ValueError(f"Multiple zones found for '{name}'") + else: + return zones[0] + + async def get_all_records(self, zone_id: str) -> List[RecordResponse]: + session = self._get_session() + async with session.get( + f"https://{ self.api_domain }/api/v1/records", params={"zone_id": zone_id} + ) as resp: + resp.raise_for_status() + data = await resp.json() + return [RecordResponse(**record) for record in data["records"]] + + async def get_record(self, record_id: str) -> RecordResponse: + session = self._get_session() + async with session.get( + f"https://{ self.api_domain }/api/v1/records/{record_id}" + ) as resp: + resp.raise_for_status() + return RecordResponse(**(await resp.json())["record"]) + + async def get_records_by_name( + self, zone_id: str, name: str + ) -> List[RecordResponse]: + records = await self.get_all_records(zone_id) + return [record for record in records if record.name == name] + + async def delete_record(self, record_id: str) -> None: + session = self._get_session() + async with session.delete( + f"https://{ self.api_domain }/api/v1/records/{record_id}" + ) as resp: + resp.raise_for_status() + return None + + async def create_record( + self, + zone_id: str, + name: str, + record_type: str, + value: str, + ttl: Optional[int] = None, + ): + session = self._get_session() + data: Dict[str, str | int] = { + "name": name, + "type": record_type, + "value": value, + "zone_id": zone_id, + } + if ttl is not None: + data["ttl"] = ttl + async with session.post( + f"https://{ self.api_domain }/api/v1/records", json=data + ) as resp: + resp.raise_for_status() + return RecordResponse(**(await resp.json())["record"]) diff --git a/desmo_api/models.py b/desmo_api/models.py new file mode 100644 index 0000000..a4df14c --- /dev/null +++ b/desmo_api/models.py @@ -0,0 +1,32 @@ +from pydantic import BaseModel, ValidationError, validator +from typing import Literal + + +class JailInfo(BaseModel): + name: str + state: str + ip: str + host: str + + +class FullJailInfo(JailInfo): + packages: list[str] = [] + commands: list[str] = [] + + +class CreateJailRequest(BaseModel): + name: str + packages: list[str] = [] + commands: list[str] = [] + + @validator("name") + def server_name_validator(cls, v: str): + if len(v) < 4 or len(v) > 32: + raise ValueError("server_name must be between 4 and 32 characters") + if not v.isascii(): + raise ValueError("server_name must be ascii") + if not v.isalnum(): + raise ValueError("server_name must be alphanumeric") + if v[0].isdigit(): + raise ValueError("server_name must not start with a number") + return v |
