"""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())