First API implementation
This commit is contained in:
58
database.py
58
database.py
@@ -84,20 +84,33 @@ class Database:
|
||||
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:
|
||||
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()
|
||||
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:
|
||||
@@ -115,6 +128,12 @@ class Database:
|
||||
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
|
||||
#
|
||||
@@ -127,6 +146,7 @@ class Database:
|
||||
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."""
|
||||
@@ -137,6 +157,7 @@ class Database:
|
||||
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(
|
||||
@@ -346,4 +367,21 @@ class Database:
|
||||
if ip in self._bans:
|
||||
del self._bans[ip]
|
||||
except:
|
||||
self._logger.error(traceback.format_exc())
|
||||
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]
|
||||
Reference in New Issue
Block a user