diff --git a/README.md b/README.md index 15e8783..f54c16b 100644 --- a/README.md +++ b/README.md @@ -51,7 +51,9 @@ python main.py сервиса - всегда будет использовано имя, указанное вами) - `!deauth ` - удалить токен - `!name [SERVICE_NAME]` - указать (или удалить) имя сервиса для токена -- `!leave` - покинуть комнату +- `!ban [REASON]` - забанить указанный IP на указанное число секунд (можно указать причину) +- `!unban ` - разбанить указанный IP адрес +- `!bans` - получить список забаненных IP адресов ## Как работает веб-сервер diff --git a/bot_callbacks.py b/bot_callbacks.py index 87cbf63..8676e14 100644 --- a/bot_callbacks.py +++ b/bot_callbacks.py @@ -173,9 +173,61 @@ async def _on_cmd_name(room: MatrixRoom, args: list[str]) -> None: # respond await _bot.send_text_to_room(room.room_id, response, is_html=True) -async def _on_cmd_leave(room: MatrixRoom, args: list[str]) -> None: - """!leave""" - await _bot.send_text_to_room(room.room_id, "Команда не реализована", is_html=True) +async def _on_cmd_ban(room: MatrixRoom, args: list[str]) -> None: + """!ban""" + # check arguments + if len(args) < 2: + error = "Формат: !ban <IP> <SECONDS> [REASON]" + await _bot.send_text_to_room(room.room_id, error, is_html=True) + return + # get args + try: + ip = args[0] + duration = float(args[1]) + reason = " ".join(args[2:]) if args[2:] else "Manual ban" + except: + error = "Возникла ошибка. Наверняка неправильно указаны секунды." + await _bot.send_text_to_room(room.room_id, error, is_html=True) + return + # ban + try: + await _db.ban_create(ip, time.time() + duration, reason) + await _bot.send_text_to_room(room.room_id, "IP адрес заблокирован", is_html=True) + except: + await _bot.send_text_to_room(room.room_id, "Возникла ошибка", is_html=True) + +async def _on_cmd_unban(room: MatrixRoom, args: list[str]) -> None: + """!unban""" + # check arguments + if len(args) != 1: + error = "Формат: !ban <IP>" + await _bot.send_text_to_room(room.room_id, error, is_html=True) + return + # unban + try: + await _db.ban_delete(args[0]) + await _bot.send_text_to_room(room.room_id, "IP адрес разблокирован (если он был заблокирован)", is_html=True) + except: + await _bot.send_text_to_room(room.room_id, "Возникла ошибка", is_html=True) + +async def _on_cmd_bans(room: MatrixRoom, args: list[str]) -> None: + """!bans""" + # list + bans = _db.ban_get_all() + if not bans: + await _bot.send_text_to_room(room.room_id, "Нет заблокированных IP адресов", is_html=True) + return + await _bot.send_text_to_room(room.room_id, f"Статус бана 127.0.0.1: {"БАН" if _db.ban_check("127.0.0.1") else "небан"}", is_html=True) + # create the response + result = "Список заблокированных IP
    " + for ban in bans: + result += f"
  • {html.escape(ban["ip"])}
      " + result += f"
    • Действует до: {util.date_to_text(ban["expires_at"])}
    • " + result += f"
    • Причина: {html.escape(ban["reason"])}
    • " + result += "
  • " + result += "
" + await _bot.send_text_to_room(room.room_id, result, is_html=True) + _COMMANDS = { ("help", "h", "?"): (_on_cmd_help, "Получить справку"), @@ -184,7 +236,9 @@ _COMMANDS = { ("auth", "create", "a", "c"): (_on_cmd_auth, "Создать новый токен"), ("deauth", "delete", "d"): (_on_cmd_deauth, "Удалить существующий токен"), ("name", "n"): (_on_cmd_name, "Задать (или удалить) имя для токена"), - ("leave", "l"): (_on_cmd_leave, "Покинуть комнату, в которой получена эта команда"), + ("ban", "b"): (_on_cmd_ban, "Забанить IP адрес"), + ("unban", "u"): (_on_cmd_unban, "Разбанить IP адрес"), + ("bans", "l"): (_on_cmd_bans, "Получить список забаненных IP адресов"), } # diff --git a/database.py b/database.py index c9e3d9e..7c38052 100644 --- a/database.py +++ b/database.py @@ -1,5 +1,6 @@ """This module implements database operations.""" +import asyncio import logging import traceback import time @@ -17,6 +18,9 @@ class Database: # # PRIVATE # + BACKGROUND_ROUTINE_PERIOD = 60 + + @staticmethod async def _setup_tables(conn: Connection) -> None: SETUP_SQL_SCRIPT = """ @@ -32,12 +36,6 @@ class Database: 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, @@ -70,14 +68,65 @@ class Database: 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) + ) + # delete old bans + try: + if self._connection is not None: + statement = "DELETE FROM bans WHERE expires_at <= ? RETURNING ip" + async with self._connection.execute(statement, (time.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()) + 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 + # # 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 = {} async def connect(self) -> bool: """Connect to the database. Returns False on failure.""" @@ -87,6 +136,12 @@ class Database: self._connection = await aiosqlite.connect(self._path) self._connection.row_factory = Row await self._setup_tables(self._connection) + self._bans = {} + 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 @@ -96,8 +151,13 @@ class Database: """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: @@ -245,5 +305,45 @@ class Database: 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 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()) \ No newline at end of file