Removed !leave, added IP ban system
This commit is contained in:
112
database.py
112
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())
|
||||
Reference in New Issue
Block a user