Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 60 additions & 29 deletions src/cmds/automation/scheduled_tasks.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
import asyncio
import logging
from collections.abc import Coroutine
from datetime import datetime, timedelta

from discord.ext import commands, tasks
from sqlalchemy import select
from sqlalchemy.exc import NoResultFound

from src import settings
from src.bot import Bot
Expand All @@ -20,25 +22,52 @@ class ScheduledTasks(commands.Cog):

def __init__(self, bot: Bot):
self.bot = bot
self._pending_tasks: set[asyncio.Task] = set()
self.all_tasks.start()

def _track_task(self, coro: Coroutine) -> None:
"""Retain a background task so it is not garbage-collected while it sleeps."""
task = asyncio.create_task(coro)
self._pending_tasks.add(task)
task.add_done_callback(self._pending_tasks.discard)

async def _run_scheduled(self, action: Coroutine, run_at: datetime, kind: str) -> None:
"""Run a ban or mute removal at `run_at` without blocking the scheduler loop."""
try:
await schedule(action, run_at=run_at)
except NoResultFound:
logger.info("Scheduled %s skipped; the record was already cleared.", kind)
except Exception:
logger.exception("Scheduled %s failed.", kind)

@staticmethod
def _due_before_epoch() -> int:
"""
Return the unix-second cutoff for records due within the next minute.

unban_time and unmute_time are stored as seconds. A millisecond cutoff
would select every row, including bans years in the future.
"""
return int((datetime.now() + timedelta(minutes=1)).timestamp())

@tasks.loop(minutes=1)
async def all_tasks(self) -> None:
"""Gathers all scheduled tasks."""
"""Schedule due unbans and unmutes without waiting for them to finish."""
logger.debug("Gathering scheduled tasks...")
await self.auto_unban()
await self.auto_unmute()
# await asyncio.gather(self.auto_unmute())
try:
await self.auto_unban()
await self.auto_unmute()
except Exception:
logger.exception("Failed while scheduling unbans or unmutes.")
logger.debug("Scheduling completed.")

async def auto_unban(self) -> None:
"""Task to automatically unban members."""
unban_tasks = []
unban_time = datetime.timestamp(datetime.now() + timedelta(minutes=1)) * 1000
logger.debug(f"Checking for bans to remove until {unban_time}.")
"""Schedule removal of bans that expire within the next minute."""
due_before = self._due_before_epoch()
logger.debug(f"Checking for bans to remove until {due_before}.")
async with AsyncSessionLocal() as session:
result = await session.scalars(
select(Ban).filter(Ban.unbanned.is_(False)).filter(Ban.unban_time <= unban_time)
select(Ban).filter(Ban.unbanned.is_(False)).filter(Ban.unban_time <= due_before)
)
bans = result.all()
logger.debug(f"Got {len(bans)} bans from DB.")
Expand All @@ -51,27 +80,27 @@ async def auto_unban(self) -> None:
continue

for ban in bans:
run_at = datetime.fromtimestamp(ban.unban_time)
logger.debug(
f"Got user_id: {ban.user_id} and unban timestamp: {run_at} from DB."
)
try:
run_at = datetime.fromtimestamp(ban.unban_time)
except (OverflowError, OSError, ValueError):
logger.exception(
"Invalid unban_time %s for user_id %s.", ban.unban_time, ban.user_id
)
continue
logger.debug(f"Got user_id: {ban.user_id} and unban timestamp: {run_at} from DB.")
member = await self.bot.get_member_or_user(guild, ban.user_id)
if not member:
logger.info(f"Member with id: {ban.user_id} not found.")
continue
unban_task = schedule(unban_member(guild, member), run_at=run_at)
unban_tasks.append(unban_task)
self._track_task(self._run_scheduled(unban_member(guild, member), run_at, "unban"))
logger.info(f"Scheduled unban task for user_id {ban.user_id} at {run_at}.")

await asyncio.gather(*unban_tasks)

async def auto_unmute(self) -> None:
"""Task to automatically unmute members."""
unmute_tasks = []
unmute_time = datetime.timestamp(datetime.now() + timedelta(minutes=1)) * 1000
logger.debug(f"Checking for mutes to remove until {unmute_time}.")
"""Schedule removal of mutes that expire within the next minute."""
due_before = self._due_before_epoch()
logger.debug(f"Checking for mutes to remove until {due_before}.")
async with AsyncSessionLocal() as session:
result = await session.scalars(select(Mute).filter(Mute.unmute_time <= unmute_time))
result = await session.scalars(select(Mute).filter(Mute.unmute_time <= due_before))
mutes = result.all()
logger.debug(f"Got {len(mutes)} mutes from DB.")

Expand All @@ -82,7 +111,13 @@ async def auto_unmute(self) -> None:
continue

for mute in mutes:
run_at = datetime.fromtimestamp(mute.unmute_time)
try:
run_at = datetime.fromtimestamp(mute.unmute_time)
except (OverflowError, OSError, ValueError):
logger.exception(
"Invalid unmute_time %s for user_id %s.", mute.unmute_time, mute.user_id
)
continue
logger.debug(
"Got user_id: {user_id} and unmute timestamp: {unmute_ts} from DB.".format(
user_id=mute.user_id, unmute_ts=run_at
Expand All @@ -92,12 +127,8 @@ async def auto_unmute(self) -> None:
if not member:
logger.info(f"Member with id: {mute.user_id} not found.")
continue
unmute_task = schedule(unmute_member(guild, member), run_at=run_at)
unmute_tasks.append(unmute_task)
logger.info(f"Scheduled unban task for user_id {mute.user_id} at {str(run_at)}.")


await asyncio.gather(*unmute_tasks)
self._track_task(self._run_scheduled(unmute_member(guild, member), run_at, "unmute"))
logger.info(f"Scheduled unmute task for user_id {mute.user_id} at {run_at}.")


def setup(bot: Bot) -> None:
Expand Down
27 changes: 24 additions & 3 deletions src/cmds/core/mute.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,12 @@
import asyncio
import logging
from datetime import datetime

from discord import ApplicationContext, Interaction, WebhookMessage, slash_command, Member
from discord import ApplicationContext, Guild, Interaction, Member, WebhookMessage, slash_command
from discord.errors import Forbidden
from discord.ext import commands
from discord.ext.commands import has_any_role
from sqlalchemy.exc import NoResultFound

from src.bot import Bot
from src.core import settings
Expand All @@ -14,12 +17,30 @@
from src.helpers.duration import validate_duration
from src.helpers.schedule import schedule

logger = logging.getLogger(__name__)


class MuteCog(commands.Cog):
"""Mute related commands."""

def __init__(self, bot: Bot):
self.bot = bot
self._pending_tasks: set[asyncio.Task] = set()

def _schedule_unmute(self, guild: Guild, member: Member, run_at: datetime) -> None:
"""Unmute `member` at `run_at`, keeping a reference so the task cannot be collected."""
task = self.bot.loop.create_task(self._unmute_when_due(guild, member, run_at))
self._pending_tasks.add(task)
task.add_done_callback(self._pending_tasks.discard)

async def _unmute_when_due(self, guild: Guild, member: Member, run_at: datetime) -> None:
"""Remove the mute when its duration elapses."""
try:
await schedule(unmute_member(guild, member), run_at=run_at)
except NoResultFound:
logger.info("Mute for user_id %s was already removed.", member.id)
except Exception:
logger.exception("Failed to unmute user_id %s.", member.id)

@slash_command(
guild_ids=settings.guild_ids,
Expand Down Expand Up @@ -56,8 +77,8 @@ async def mute(
if isinstance(member, Member):
role = ctx.guild.get_role(settings.roles.MUTED)
await member.add_roles(role)
timestamp=datetime.fromtimestamp(dur)
self.bot.loop.create_task(schedule(unmute_member(ctx.guild, member), run_at=timestamp))
timestamp = datetime.fromtimestamp(dur)
self._schedule_unmute(ctx.guild, member, timestamp)
await member.timeout(timestamp, reason=reason if reason else "Time to shush, innit?")
try:
await member.send(f"You have been muted for {duration}. Reason:\n>>> {reason}")
Expand Down
175 changes: 175 additions & 0 deletions tests/src/cmds/automation/test_scheduled_tasks.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,175 @@
import asyncio
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest import mock

import pytest
from sqlalchemy.exc import NoResultFound
from sqlalchemy.sql.elements import BindParameter

from src.cmds.automation import scheduled_tasks
from tests import helpers


def _bound_values(stmt):
values = []
for bind in stmt.compile().binds.values():
if isinstance(bind, BindParameter) and isinstance(bind.value, (int, float)):
values.append(bind.value)
return values


class _Session:
def __init__(self, rows, column):
self.rows = rows
self.column = column
self.stmt = None

async def __aenter__(self):
return self

async def __aexit__(self, exc_type, exc, tb):
return False

async def scalars(self, stmt):
self.stmt = stmt
cutoffs = [value for value in _bound_values(stmt) if value > 10**8]
cutoff = cutoffs[-1] if cutoffs else 0
matched = [row for row in self.rows if getattr(row, self.column) <= cutoff]
result = mock.Mock()
result.all.return_value = matched
return result


def _cog(bot):
with mock.patch("discord.ext.tasks.Loop.start"):
return scheduled_tasks.ScheduledTasks(bot)


class TestScheduledTasks:
"""Unmutes must be scheduled in seconds and must not wait on long bans."""

def test_due_before_epoch_is_unix_seconds(self):
now = datetime.now().timestamp()
due_before = scheduled_tasks.ScheduledTasks._due_before_epoch()

assert now < due_before < now + 90
assert due_before < 10**11

@pytest.mark.asyncio
async def test_far_future_ban_does_not_block_scheduling(self, bot):
"""A multi-year ban must not stall the loop that also unmutes members."""
guild = helpers.MockGuild()
member = helpers.MockMember()
bot.get_guild.return_value = guild
bot.get_member_or_user = mock.AsyncMock(return_value=member)
far_future = int((datetime.now() + timedelta(weeks=500)).timestamp())
ban = SimpleNamespace(user_id=member.id, unban_time=far_future)
session = _Session([ban], "unban_time")
cog = _cog(bot)

with mock.patch.object(scheduled_tasks.settings, "guild_ids", [guild.id]), mock.patch.object(
scheduled_tasks, "AsyncSessionLocal", return_value=session
):
await asyncio.wait_for(cog.auto_unban(), timeout=1)

assert cog._pending_tasks == set()

@pytest.mark.asyncio
async def test_expired_mute_is_cleared_without_waiting_on_future_mutes(self, bot):
guild = helpers.MockGuild()
member = helpers.MockMember()
bot.get_guild.return_value = guild
bot.get_member_or_user = mock.AsyncMock(return_value=member)
now = int(datetime.now().timestamp())
expired = SimpleNamespace(user_id=member.id, unmute_time=now - 30)
still_muted = SimpleNamespace(user_id=member.id + 1, unmute_time=now + 7200)
session = _Session([expired, still_muted], "unmute_time")
cog = _cog(bot)

with mock.patch.object(scheduled_tasks.settings, "guild_ids", [guild.id]), mock.patch.object(
scheduled_tasks, "AsyncSessionLocal", return_value=session
), mock.patch.object(scheduled_tasks, "unmute_member", new_callable=mock.AsyncMock) as unmute:
await asyncio.wait_for(cog.auto_unmute(), timeout=1)
pending = list(cog._pending_tasks)
if pending:
await asyncio.wait_for(asyncio.gather(*pending), timeout=1)

unmute.assert_awaited_once()
assert session.stmt is not None
cutoff = max(value for value in _bound_values(session.stmt) if value > 10**8)
assert now < cutoff < now + 90

@pytest.mark.asyncio
async def test_due_ban_is_scheduled(self, bot):
guild = helpers.MockGuild()
member = helpers.MockMember()
bot.get_guild.return_value = guild
bot.get_member_or_user = mock.AsyncMock(return_value=member)
ban = SimpleNamespace(user_id=member.id, unban_time=int(datetime.now().timestamp()) - 15)
session = _Session([ban], "unban_time")
cog = _cog(bot)

with mock.patch.object(scheduled_tasks.settings, "guild_ids", [guild.id]), mock.patch.object(
scheduled_tasks, "AsyncSessionLocal", return_value=session
), mock.patch.object(scheduled_tasks, "unban_member", new_callable=mock.AsyncMock) as unban:
await cog.auto_unban()
await asyncio.wait_for(asyncio.gather(*cog._pending_tasks), timeout=1)

unban.assert_awaited_once()

@pytest.mark.asyncio
async def test_missing_guild_member_and_bad_timestamp_are_skipped(self, bot):
guild = helpers.MockGuild()
bot.get_guild.side_effect = lambda guild_id: guild if guild_id == guild.id else None
bot.get_member_or_user = mock.AsyncMock(return_value=None)
now = int(datetime.now().timestamp())
mute = SimpleNamespace(user_id=1, unmute_time=now - 5)
bad_mute = SimpleNamespace(user_id=2, unmute_time=10**18)
ban = SimpleNamespace(user_id=3, unban_time=now - 5)
bad_ban = SimpleNamespace(user_id=4, unban_time=10**18)
cog = _cog(bot)

def session_for(rows):
session = mock.Mock()
session.__aenter__ = mock.AsyncMock(return_value=session)
session.__aexit__ = mock.AsyncMock(return_value=False)
result = mock.Mock()
result.all.return_value = rows
session.scalars = mock.AsyncMock(return_value=result)
return session

with mock.patch.object(scheduled_tasks.settings, "guild_ids", [guild.id + 1, guild.id]), mock.patch.object(
scheduled_tasks, "AsyncSessionLocal", side_effect=[session_for([ban, bad_ban]), session_for([mute, bad_mute])]
):
await cog.auto_unban()
await cog.auto_unmute()

assert cog._pending_tasks == set()

@pytest.mark.asyncio
async def test_all_tasks_schedules_both_and_survives_errors(self, bot):
cog = _cog(bot)
cog.auto_unban = mock.AsyncMock()
cog.auto_unmute = mock.AsyncMock()
await cog.all_tasks()
cog.auto_unban.assert_awaited_once()
cog.auto_unmute.assert_awaited_once()

cog.auto_unban.side_effect = RuntimeError("database unavailable")
await cog.all_tasks()
cog.auto_unmute.assert_awaited_once()

@pytest.mark.asyncio
async def test_run_scheduled_swallows_cleared_records_and_failures(self, bot):
cog = _cog(bot)
run_at = datetime.now() - timedelta(seconds=1)

async def already_cleared():
raise NoResultFound("mute already removed")

async def failed():
raise RuntimeError("discord unavailable")

await cog._run_scheduled(already_cleared(), run_at, "unmute")
await cog._run_scheduled(failed(), run_at, "unban")
Loading
Loading