diff --git a/src/adapters/discord_bot/bot.py b/src/adapters/discord_bot/bot.py index 822713d..c62b6a6 100644 --- a/src/adapters/discord_bot/bot.py +++ b/src/adapters/discord_bot/bot.py @@ -105,7 +105,13 @@ async def on_tree_error( self, interaction: discord.Interaction, error: discord.app_commands.AppCommandError ) -> None: """Global fallback error handler for slash commands outside cogs or tree-level errors.""" - cmd_name = interaction.command.qualified_name if interaction.command else "command" + command = interaction.command + if command is not None: + has_handlers = getattr(command, "_has_any_error_handlers", None) + if callable(has_handlers) and has_handlers(): + return + + cmd_name = command.qualified_name if command else "command" await send_interaction_error( interaction, error, diff --git a/src/adapters/discord_bot/error_handler.py b/src/adapters/discord_bot/error_handler.py index ecc5830..c9dbce2 100644 --- a/src/adapters/discord_bot/error_handler.py +++ b/src/adapters/discord_bot/error_handler.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any import discord @@ -99,10 +99,13 @@ async def send_interaction_error( ): is_done = True + dismissal_target: Any = interaction if is_done and hasattr(interaction, "followup") and callable(getattr(interaction.followup, "send", None)): - res = interaction.followup.send(message, ephemeral=ephemeral) + res = interaction.followup.send(message, ephemeral=ephemeral, wait=True) if hasattr(res, "__await__"): - await res + res = await res + if res is not None: + dismissal_target = res elif hasattr(interaction, "response") and callable(getattr(interaction.response, "send_message", None)): res = interaction.response.send_message(message, ephemeral=ephemeral) if hasattr(res, "__await__"): @@ -111,7 +114,7 @@ async def send_interaction_error( if ephemeral and auto_dismiss: from src.adapters.discord_bot.menu_manager import menu_manager - menu_manager.schedule_toast_dismissal(interaction, delay=dismiss_delay) + menu_manager.schedule_toast_dismissal(dismissal_target, delay=dismiss_delay) except Exception as send_err: log.exception("Failed to send error response to Discord interaction: %s", send_err) diff --git a/tests/test_error_handling.py b/tests/test_error_handling.py index 5832846..76eb10f 100644 --- a/tests/test_error_handling.py +++ b/tests/test_error_handling.py @@ -121,10 +121,35 @@ async def test_send_interaction_error_deferred(): stale_err = StaleVersionError("Conflict") msg = await send_interaction_error(interaction, stale_err, "updating task", ephemeral=True) - interaction.followup.send.assert_awaited_once_with(msg, ephemeral=True) + interaction.followup.send.assert_awaited_once_with(msg, ephemeral=True, wait=True) assert "already modified" in msg +@pytest.mark.asyncio +async def test_send_interaction_error_deferred_schedules_followup_dismissal(monkeypatch): + """Test send_interaction_error schedules toast dismissal on the followup message, not interaction.""" + from src.adapters.discord_bot.menu_manager import menu_manager + + scheduled_targets = [] + monkeypatch.setattr( + menu_manager, "schedule_toast_dismissal", lambda target, delay=10.0: scheduled_targets.append(target) + ) + + interaction = MagicMock(spec=discord.Interaction) + interaction.response = MagicMock() + interaction.response.is_done.return_value = True + followup_msg = MagicMock(spec=discord.WebhookMessage) + interaction.followup = MagicMock() + interaction.followup.send = AsyncMock(return_value=followup_msg) + + stale_err = StaleVersionError("Conflict") + msg = await send_interaction_error(interaction, stale_err, "updating task", ephemeral=True) + + interaction.followup.send.assert_awaited_once_with(msg, ephemeral=True, wait=True) + assert len(scheduled_targets) == 1 + assert scheduled_targets[0] is followup_msg + + @pytest.mark.asyncio async def test_service_raises_typed_exceptions(services): """Verify services raise strongly typed domain exceptions.""" @@ -324,6 +349,7 @@ async def test_bot_on_tree_error_handles_missing_permissions(services, caplog): interaction = MagicMock(spec=discord.Interaction) interaction.command = MagicMock() interaction.command.qualified_name = "pm project create" + interaction.command._has_any_error_handlers.return_value = False interaction.response = MagicMock() interaction.response.is_done.return_value = False interaction.response.send_message = AsyncMock() @@ -338,3 +364,104 @@ async def test_bot_on_tree_error_handles_missing_permissions(services, caplog): assert "Manage Server" in sent_msg assert interaction.response.send_message.call_args[1].get("ephemeral") is True assert "App command check failure while executing '/pm project create'" in caplog.text + + +@pytest.mark.asyncio +async def test_bot_on_tree_error_suppresses_when_command_has_error_handler(services): + """Verify DggPmBot.on_tree_error does not send duplicate response if command has error handlers.""" + from src.adapters.discord_bot.bot import DggPmBot + + bot = DggPmBot( + task_service=services["task"], + project_service=services["project"], + squad_service=services["squad"], + ) + + interaction = MagicMock(spec=discord.Interaction) + interaction.command = MagicMock() + interaction.command.qualified_name = "pm project create" + interaction.command._has_any_error_handlers.return_value = True + interaction.response = MagicMock() + interaction.response.is_done.return_value = False + interaction.response.send_message = AsyncMock() + interaction.followup = MagicMock() + interaction.followup.send = AsyncMock() + + missing_err = discord.app_commands.MissingPermissions(["manage_guild"]) + + await bot.on_tree_error(interaction, missing_err) + + interaction.response.send_message.assert_not_called() + interaction.followup.send.assert_not_called() + + +@pytest.mark.asyncio +async def test_bot_on_tree_error_sends_followup_for_deferred_interaction_without_handlers(services): + """Verify DggPmBot.on_tree_error sends followup if deferred and command has no error handlers.""" + from src.adapters.discord_bot.bot import DggPmBot + + bot = DggPmBot( + task_service=services["task"], + project_service=services["project"], + squad_service=services["squad"], + ) + + interaction = MagicMock(spec=discord.Interaction) + interaction.command = MagicMock() + interaction.command.qualified_name = "pm project create" + interaction.command._has_any_error_handlers.return_value = False + interaction.response = MagicMock() + interaction.response.is_done.return_value = True + followup_msg = MagicMock(spec=discord.WebhookMessage) + interaction.followup = MagicMock() + interaction.followup.send = AsyncMock(return_value=followup_msg) + + missing_err = discord.app_commands.MissingPermissions(["manage_guild"]) + + await bot.on_tree_error(interaction, missing_err) + + interaction.response.send_message.assert_not_called() + interaction.followup.send.assert_awaited_once() + assert interaction.followup.send.call_args[1].get("wait") is True + + +@pytest.mark.asyncio +async def test_command_check_failure_pipeline_dispatches_single_error_response(services): + """End-to-end test simulating discord.py tree dispatch: cog handler responds and tree handler suppresses.""" + from src.adapters.discord_bot.bot import DggPmBot + + bot = DggPmBot( + task_service=services["task"], + project_service=services["project"], + squad_service=services["squad"], + ) + pm_cog = PmCog( + bot, + project_service=services["project"], + squad_service=services["squad"], + task_service=services["task"], + ) + + interaction = MagicMock(spec=discord.Interaction) + interaction.command = MagicMock() + interaction.command.qualified_name = "pm project create" + # Command in cog has error handlers + interaction.command._has_any_error_handlers.return_value = True + interaction.response = MagicMock() + interaction.response.is_done.return_value = False + interaction.response.send_message = AsyncMock() + interaction.followup = MagicMock() + interaction.followup.send = AsyncMock() + + missing_err = discord.app_commands.MissingPermissions(["manage_guild"]) + + # Step 1: discord.py invokes command error handler (in cog) + await pm_cog.cog_app_command_error(interaction, missing_err) + interaction.response.is_done.return_value = True + + # Step 2: discord.py invokes tree on_error fallback + await bot.on_tree_error(interaction, missing_err) + + # Verification: only a single response was sent, tree fallback suppressed duplicate + interaction.response.send_message.assert_awaited_once() + interaction.followup.send.assert_not_called()