diff --git a/src/cmds/automation/scheduled_tasks.py b/src/cmds/automation/scheduled_tasks.py index 9b763fc3..112fd4ae 100644 --- a/src/cmds/automation/scheduled_tasks.py +++ b/src/cmds/automation/scheduled_tasks.py @@ -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 @@ -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.") @@ -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.") @@ -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 @@ -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: diff --git a/src/cmds/core/mute.py b/src/cmds/core/mute.py index cb1087c1..e9fdf706 100644 --- a/src/cmds/core/mute.py +++ b/src/cmds/core/mute.py @@ -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 @@ -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, @@ -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}") diff --git a/tests/src/cmds/automation/test_scheduled_tasks.py b/tests/src/cmds/automation/test_scheduled_tasks.py new file mode 100644 index 00000000..7b78a4a1 --- /dev/null +++ b/tests/src/cmds/automation/test_scheduled_tasks.py @@ -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") diff --git a/tests/src/cmds/core/test_mute.py b/tests/src/cmds/core/test_mute.py index 0133c7ad..af892fd1 100644 --- a/tests/src/cmds/core/test_mute.py +++ b/tests/src/cmds/core/test_mute.py @@ -1,4 +1,26 @@ +from datetime import datetime, timedelta +from unittest import mock + +import pytest +from sqlalchemy.exc import NoResultFound + from src.cmds.core import mute +from src.cmds.core.mute import MuteCog +from tests import helpers + + +class _Session: + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + return False + + def add(self, obj): + self.obj = obj + + async def commit(self): + return None class TestMuteCog: @@ -10,3 +32,52 @@ def test_setup(self, bot): mute.setup(bot) bot.add_cog.assert_called_once() + + @pytest.mark.asyncio + async def test_mute_keeps_a_reference_to_the_unmute_task(self, bot, ctx): + member = helpers.MockMember(bot=False) + member.timeout = mock.AsyncMock() + member.send = mock.AsyncMock() + ctx.user = helpers.MockMember() + bot.get_member_or_user = mock.AsyncMock(return_value=member) + cog = MuteCog(bot) + unmute_at = int(datetime.now().timestamp()) + 60 + + with mock.patch("src.cmds.core.mute.member_is_staff", return_value=False), mock.patch( + "src.cmds.core.mute.validate_duration", return_value=(unmute_at, "") + ), mock.patch("src.cmds.core.mute.AsyncSessionLocal", return_value=_Session()): + await cog.mute.callback(cog, ctx, member, "1m", "testing") + + assert len(cog._pending_tasks) == 1 + member.timeout.assert_awaited_once() + + @pytest.mark.asyncio + async def test_unmute_when_due_removes_the_mute(self, bot): + cog = MuteCog(bot) + member = helpers.MockMember() + guild = helpers.MockGuild() + run_at = datetime.now() - timedelta(seconds=1) + + with mock.patch("src.cmds.core.mute.unmute_member", new_callable=mock.AsyncMock) as unmute: + await cog._unmute_when_due(guild, member, run_at) + + unmute.assert_awaited_once() + + @pytest.mark.asyncio + async def test_unmute_when_due_ignores_a_cleared_record(self, bot): + cog = MuteCog(bot) + run_at = datetime.now() - timedelta(seconds=1) + + with mock.patch("src.cmds.core.mute.unmute_member", side_effect=NoResultFound("already removed")): + await cog._unmute_when_due(helpers.MockGuild(), helpers.MockMember(), run_at) + + @pytest.mark.asyncio + async def test_unmute_when_due_logs_unexpected_failures(self, bot): + cog = MuteCog(bot) + run_at = datetime.now() - timedelta(seconds=1) + + async def explode(*args, **kwargs): + raise RuntimeError("discord unavailable") + + with mock.patch("src.cmds.core.mute.unmute_member", side_effect=explode): + await cog._unmute_when_due(helpers.MockGuild(), helpers.MockMember(), run_at)