First API implementation

This commit is contained in:
Nikita Tyukalov, ASUS, Linux
2026-08-29 17:24:39 +03:00
parent 28ef360d5f
commit 8094aaab10
6 changed files with 278 additions and 31 deletions

View File

@@ -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]