Files
2026-matrix-csonac/database.py
nikita 4194f3da48 mab update, !short command, last_access fix
- Updated mab to v0.5.0
- Added `!short` command to add named tokens
- Fixed `last_access_at` field not being updated
2026-09-13 00:12:00 +03:00

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]