- Updated mab to v0.5.0 - Added `!short` command to add named tokens - Fixed `last_access_at` field not being updated
404 lines
14 KiB
Python
404 lines
14 KiB
Python
"""This module implements database operations."""
|
|
|
|
import asyncio
|
|
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
|
|
#
|
|
BACKGROUND_ROUTINE_PERIOD = 60
|
|
|
|
|
|
@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 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
|
|
|
|
async def _background_routine(self, stop_event: asyncio.Event) -> None:
|
|
"""This routine perform routine tasks."""
|
|
stop_task = asyncio.create_task(stop_event.wait())
|
|
wait_task = asyncio.create_task(asyncio.sleep(0))
|
|
self._logger.debug("Started background worker")
|
|
while True:
|
|
done, _ = await asyncio.wait(
|
|
[stop_task, wait_task],
|
|
return_when=asyncio.FIRST_COMPLETED
|
|
)
|
|
if stop_task in done:
|
|
wait_task.cancel()
|
|
break
|
|
wait_task = asyncio.create_task(
|
|
asyncio.sleep(self.BACKGROUND_ROUTINE_PERIOD)
|
|
)
|
|
if self._connection is None:
|
|
continue
|
|
# get current time
|
|
current_time = time.time()
|
|
# delete old bans
|
|
try:
|
|
statement = "DELETE FROM bans WHERE expires_at <= ? RETURNING ip"
|
|
async with self._connection.execute(statement, (current_time,)) as cursor:
|
|
async for row in cursor:
|
|
ip = row["ip"]
|
|
if ip in self._bans:
|
|
del self._bans[ip]
|
|
self._logger.debug(f"IP {ip} is not banned anymore")
|
|
await self._connection.commit()
|
|
self._logger.debug("Performed banned IPs cleanup")
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
# delete expired fails
|
|
try:
|
|
for ip in dict(self._fails):
|
|
self._fails[ip] = list(
|
|
filter(lambda x: x > current_time, self._fails[ip])
|
|
)
|
|
if not self._fails[ip]:
|
|
del self._fails[ip]
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
self._logger.debug("Stopped background worker")
|
|
|
|
async def _load_bans_from_database(self) -> None:
|
|
"""Loads bans information from database. Must be called when database connection is established."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
data = await self._select_by_anded_kwargs(self._connection, "bans")
|
|
for d in data:
|
|
self._bans[d["ip"]] = {
|
|
"expires_at": d["expires_at"],
|
|
"reason": d["reason"]
|
|
}
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
return
|
|
|
|
def _get_fails_for_ip(self, ip) -> int:
|
|
if ip not in self._fails:
|
|
return 0
|
|
t = time.time()
|
|
return sum([(1 if expires_at > t else 0) for expires_at in self._fails[ip]])
|
|
|
|
#
|
|
# PUBLIC
|
|
#
|
|
def __init__(self, path: Path):
|
|
self._path = path
|
|
self._logger = logging.getLogger("database")
|
|
self._logger.setLevel(logging.DEBUG)
|
|
|
|
self._connection: Connection | None = None
|
|
self._background_stop_event: asyncio.Event | None = None
|
|
self._background_task: asyncio.Task | None = None
|
|
self._bans: dict = {}
|
|
self._fails: dict[str, list[float]] = {}
|
|
|
|
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)
|
|
self._bans = {}
|
|
self._fails = {}
|
|
await self._load_bans_from_database()
|
|
self._background_stop_event = asyncio.Event()
|
|
self._background_task = asyncio.create_task(
|
|
self._background_routine(self._background_stop_event)
|
|
)
|
|
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
|
|
self._background_stop_event.set() # type: ignore
|
|
await self._background_task # type: ignore
|
|
|
|
await self._connection.close()
|
|
self._connection = None
|
|
self._background_task = None
|
|
self._background_stop_event = 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, code: str | None = None) -> ObjectToken:
|
|
"""Create a token."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
if code is None:
|
|
code = get_hash(str(time.time()).encode() + os.urandom(64))
|
|
else:
|
|
if len(code) < 2 or len(code) > 16:
|
|
raise RuntimeError("Token must be 2-16 symbols long")
|
|
if not all(ord(c) < 128 and (c.islower() or c.isdigit()) for c in code):
|
|
raise RuntimeError("Only digits and ASCII lowercase letters are allowed")
|
|
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())
|
|
|
|
async def token_set_last_access(self, code: str, timestamp: float) -> None:
|
|
"""Set last_access_at for the token."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
statement = "UPDATE tokens SET last_access_at=? WHERE code=?"
|
|
await self._connection.execute(statement, (timestamp, code))
|
|
await self._connection.commit()
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
|
|
|
|
async def ban_create(self, ip: str, expires_at: float, reason: str) -> None:
|
|
"""Save information about banned IP address. Replaces existing IPs."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
statement = "INSERT OR REPLACE INTO bans (ip, expires_at, reason) VALUES (?, ?, ?)"
|
|
await self._connection.execute(statement, (ip, expires_at, reason))
|
|
await self._connection.commit()
|
|
self._bans[ip] = {
|
|
"expires_at": expires_at,
|
|
"reason": reason
|
|
}
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
|
|
def ban_get_all(self) -> list[dict]:
|
|
"""Get information about banned IP addresses."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
return [{"ip": k, **self._bans[k]} for k in self._bans]
|
|
|
|
def ban_check(self, ip: str) -> bool:
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
return ip in self._bans
|
|
|
|
async def ban_delete(self, ip: str) -> None:
|
|
"""Unbans specified IP address."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
try:
|
|
statement = "DELETE FROM bans WHERE ip = ?"
|
|
await self._connection.execute(statement, (ip,))
|
|
await self._connection.commit()
|
|
if ip in self._bans:
|
|
del self._bans[ip]
|
|
except:
|
|
self._logger.error(traceback.format_exc())
|
|
|
|
|
|
def fail_create(self, ip: str, expires_at: float) -> int:
|
|
"""Adds failed access attempt. Returns count of failed attempts for IP."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
if ip not in self._fails:
|
|
self._fails[ip] = []
|
|
self._fails[ip].append(expires_at)
|
|
return self._get_fails_for_ip(ip)
|
|
|
|
def fail_clear(self, ip: str) -> None:
|
|
"""Clears failed attempts counter."""
|
|
if self._connection is None:
|
|
raise RuntimeError("Not connected to the database")
|
|
if ip in self._fails:
|
|
del self._fails[ip] |