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
125
126
127
128
129
130
131
132
133
134
135
136
137
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,
)
|