Basic database function implemented
This commit is contained in:
249
database.py
249
database.py
@@ -0,0 +1,249 @@
|
||||
"""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())
|
||||
Reference in New Issue
Block a user