249 lines
8.0 KiB
Python
249 lines
8.0 KiB
Python
"""This module implements database operations."""
|
|
|
|
import logging
|
|
import traceback
|
|
import time
|
|
import os
|
|
from pathlib import Path
|
|
from typing import Any
|
|
import aiosqlite
|
|
from aiosqlite import Connection, Row
|
|
|
|
from datatypes import *
|
|
from util import get_hash
|
|
|
|
|
|
class Database:
|
|
#
|
|
# PRIVATE
|
|
#
|
|
@staticmethod
|
|
async def _setup_tables(conn: Connection) -> None:
|
|
SETUP_SQL_SCRIPT = """
|
|
CREATE TABLE IF NOT EXISTS rooms (
|
|
code TEXT PRIMARY KEY,
|
|
matrix_id TEXT NOT NULL UNIQUE
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS tokens (
|
|
code TEXT PRIMARY KEY,
|
|
name TEXT,
|
|
created_at REAL,
|
|
last_access_at REAL
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS fails (
|
|
ip TEXT PRIMARY KEY,
|
|
timestamp INTEGER,
|
|
score INTEGER
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS bans (
|
|
ip TEXT PRIMARY KEY,
|
|
expires_at INTEGER,
|
|
reason TEXT
|
|
);
|
|
"""
|
|
await conn.executescript(SETUP_SQL_SCRIPT)
|
|
await conn.commit()
|
|
|
|
@staticmethod
|
|
async def _select_by_anded_kwargs(conn: Connection, table_name: str, **kwargs) -> list[dict[str, Any]]:
|
|
# prepare keys and values
|
|
keys = tuple(k for k in kwargs)
|
|
values = tuple(kwargs[k] for k in keys)
|
|
# prepare the statement
|
|
statement = f"SELECT * FROM {table_name}"
|
|
# add `WHERE` part if there kwargs
|
|
if keys:
|
|
statement += " WHERE "
|
|
statement += " AND ".join([f"{k}=?" for k in keys])
|
|
else:
|
|
values = None
|
|
# execute and enumerate
|
|
result = []
|
|
async with conn.execute(statement, values) as cursor:
|
|
async for row in cursor:
|
|
entry = {}
|
|
for k in row.keys():
|
|
entry[k] = row[k]
|
|
result.append(entry)
|
|
return result
|
|
|
|
#
|
|
# PUBLIC
|
|
#
|
|
def __init__(self, path: Path):
|
|
self._path = path
|
|
self._logger = logging.getLogger("database")
|
|
|
|
self._connection: Connection | None = None
|
|
|
|
async def connect(self) -> bool:
|
|
"""Connect to the database. Returns False on failure."""
|
|
if self._connection is not None:
|
|
return False
|
|
try:
|
|
self._connection = await aiosqlite.connect(self._path)
|
|
self._connection.row_factory = Row
|
|
await self._setup_tables(self._connection)
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
return False
|
|
return True
|
|
|
|
async def disconnect(self) -> None:
|
|
"""Disconnect from the database."""
|
|
if self._connection is None:
|
|
return
|
|
await self._connection.close()
|
|
self._connection = None
|
|
|
|
|
|
async def room_create(self, matrix_id: str) -> ObjectRoom | None:
|
|
"""Create a room with specified matrix_id. Returns None if the room exists."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
# room exists
|
|
if await self.room_get(matrix_id) is not None:
|
|
return None
|
|
# create the code
|
|
code = None
|
|
string_to_hash = matrix_id
|
|
while code is None or (await self.room_get(code)) is not None:
|
|
code = get_hash(string_to_hash)[-10:]
|
|
string_to_hash += "A"
|
|
# code must always start with q
|
|
code = f"q{code}"
|
|
# add
|
|
try:
|
|
statement = "INSERT INTO rooms (code, matrix_id) VALUES (?, ?)"
|
|
await self._connection.execute(statement, (code, matrix_id))
|
|
await self._connection.commit()
|
|
self._logger.info(f"Added new room with code {code}")
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
return None
|
|
# return the added object
|
|
return ObjectRoom(
|
|
code=code,
|
|
matrix_id=matrix_id
|
|
)
|
|
|
|
async def room_get_all(self) -> list[ObjectRoom]:
|
|
"""Get room by code."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
rooms = await self._select_by_anded_kwargs(
|
|
self._connection,
|
|
"rooms"
|
|
)
|
|
result = []
|
|
for r in rooms:
|
|
result.append(ObjectRoom(
|
|
code=r["code"],
|
|
matrix_id=r["matrix_id"]
|
|
))
|
|
return result
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
return []
|
|
|
|
async def room_get(self, identifier: str) -> ObjectRoom | None:
|
|
"""Get room by identifier (either `code` or `matrix_id`)."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
kwargs = {}
|
|
kwargs["code" if identifier[0] == "q" else "matrix_id"] = identifier
|
|
rooms = await self._select_by_anded_kwargs(
|
|
self._connection,
|
|
"rooms",
|
|
**kwargs
|
|
)
|
|
if not rooms:
|
|
self._logger.debug(f"Room `{identifier}` is not found")
|
|
return None
|
|
return ObjectRoom(**rooms[0])
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
return None
|
|
|
|
|
|
async def token_create(self) -> ObjectToken:
|
|
"""Create a token."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
code = get_hash(str(time.time()).encode() + os.urandom(64))
|
|
create_time = time.time()
|
|
statement = """
|
|
INSERT INTO tokens (code, name, created_at, last_access_at)
|
|
VALUES (?, ?, ?, ?)
|
|
"""
|
|
result = ObjectToken(
|
|
code=code,
|
|
name=None,
|
|
created_at=create_time,
|
|
last_access_at=create_time
|
|
)
|
|
await self._connection.execute(
|
|
statement,
|
|
(result.code, result.name, result.created_at, result.last_access_at)
|
|
)
|
|
await self._connection.commit()
|
|
return result
|
|
|
|
async def token_delete(self, code: str) -> None:
|
|
"""Delete a token."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
await self._connection.execute("DELETE FROM tokens WHERE code=?", (code,))
|
|
await self._connection.commit()
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
|
|
async def token_get(self, code: str) -> ObjectToken | None:
|
|
"""Get token by its code."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
result = await self._select_by_anded_kwargs(
|
|
self._connection,
|
|
"tokens",
|
|
code=code
|
|
)
|
|
if not result:
|
|
return None
|
|
result = result[0]
|
|
return ObjectToken(**result)
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
return None
|
|
|
|
async def token_get_all(self) -> list[ObjectToken]:
|
|
"""Get list of all tokens."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
result = await self._select_by_anded_kwargs(
|
|
self._connection,
|
|
"tokens"
|
|
)
|
|
result = [ObjectToken(**r) for r in result]
|
|
return result
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
return []
|
|
|
|
async def token_set_name(self, code: str, name: str | None) -> None:
|
|
"""Set name for the token."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
statement = "UPDATE tokens SET name=? WHERE code=?"
|
|
await self._connection.execute(statement, (name, code))
|
|
await self._connection.commit()
|
|
except:
|
|
self._logger.error(traceback.format_exc()) |