diff --git a/astrbot/core/platform/sources/discord/discord_platform_event.py b/astrbot/core/platform/sources/discord/discord_platform_event.py index d622b724d8..65eeb6a262 100644 --- a/astrbot/core/platform/sources/discord/discord_platform_event.py +++ b/astrbot/core/platform/sources/discord/discord_platform_event.py @@ -3,6 +3,7 @@ from io import BytesIO from pathlib import Path from typing import cast +from urllib.parse import unquote, urlparse import discord from discord.types.interactions import ComponentInteractionData @@ -16,6 +17,7 @@ Plain, Record, Reply, + Video, ) from astrbot.api.platform import ( AstrBotMessage, @@ -386,6 +388,22 @@ async def _parse_to_discord( describe_media_ref(audio_ref), exc_info=True, ) + elif isinstance(i, Video): + try: + path = Path(await i.convert_to_file_path()) + filename = path.name + if i.file.startswith(("http://", "https://")): + url_path = unquote(urlparse(i.file).path).replace("\\", "/") + url_filename = url_path.rsplit("/", 1)[-1] + if url_filename not in ("", ".", ".."): + filename = url_filename + files.append(discord.File(path, filename=filename)) + except Exception: + logger.warning( + "[Discord] Failed to process video: %s", + describe_media_ref(i.file), + exc_info=True, + ) elif isinstance(i, File): try: file_path_str = await i.get_file() diff --git a/tests/test_discord_adapter.py b/tests/test_discord_adapter.py index b0413965da..a0d38ebd2e 100644 --- a/tests/test_discord_adapter.py +++ b/tests/test_discord_adapter.py @@ -1,11 +1,11 @@ import base64 -from io import BytesIO +from io import BufferedReader, BytesIO from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, MagicMock import pytest -from astrbot.api.message_components import Image, Record +from astrbot.api.message_components import File, Image, Plain, Record, Video from astrbot.api.platform import Group, MessageType from astrbot.core.message.message_event_result import MessageChain from astrbot.core.platform.sources.discord import ( @@ -26,6 +26,89 @@ _WAV_PATH = "/tmp/discord_voice.wav" +@pytest.mark.asyncio +@pytest.mark.parametrize("component_type", ["video", "video_url", "file"]) +@pytest.mark.parametrize("with_text", [False, True]) +async def test_discord_sends_video_attachment( + tmp_path, monkeypatch, component_type, with_text +): + video_path = tmp_path / ( + "media_video_test.mp4" if component_type == "video_url" else "repro.mp4" + ) + video_bytes = b"video attachment payload" + video_path.write_bytes(video_bytes) + if component_type == "video": + component = Video.fromFileSystem(str(video_path)) + elif component_type == "video_url": + component = Video.fromURL("https://example.com/repro.mp4") + monkeypatch.setattr( + Video, + "convert_to_file_path", + AsyncMock(return_value=str(video_path)), + ) + else: + component = File(name=video_path.name, file=str(video_path)) + + channel = MagicMock(spec=discord_platform_event.discord.abc.Messageable) + channel.send = AsyncMock() + event = DiscordPlatformEvent.__new__(DiscordPlatformEvent) + event.interaction_followup_webhook = None + event._get_channel = AsyncMock(return_value=channel) + monkeypatch.setattr(discord_platform_event.AstrMessageEvent, "send", AsyncMock()) + chain = [Plain("Video test"), component] if with_text else [component] + + await event.send(MessageChain(chain=chain)) + + channel.send.assert_awaited_once() + kwargs = channel.send.call_args.kwargs + assert kwargs.get("content", "") == ("Video test" if with_text else "") + assert len(kwargs["files"]) == 1 + attachment = kwargs["files"][0] + try: + assert attachment.filename == "repro.mp4" + if component_type != "file": + assert isinstance(attachment.fp, BufferedReader) + assert attachment.fp.tell() == 0 + assert attachment.fp.read() == video_bytes + finally: + attachment.close() + if component_type != "file": + assert attachment.fp.closed + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("url", "expected_name"), + [ + ("https://example.com/clips/demo.mp4", "demo.mp4"), + ("https://example.com/demo.mp4?token=test#preview", "demo.mp4"), + ("https://example.com/%E6%B5%8B%E8%AF%95%20video.mp4", "测试 video.mp4"), + ("https://example.com", "media_video_test.mp4"), + ("https://example.com/clips/?token=test", "media_video_test.mp4"), + ], +) +async def test_discord_video_url_attachment_filename( + tmp_path, monkeypatch, url, expected_name +): + video_path = tmp_path / "media_video_test.mp4" + video_path.write_bytes(b"video attachment payload") + monkeypatch.setattr( + Video, "convert_to_file_path", AsyncMock(return_value=str(video_path)) + ) + event = DiscordPlatformEvent.__new__(DiscordPlatformEvent) + + _, files, _, _, _ = await event._parse_to_discord( + MessageChain(chain=[Video.fromURL(url)]) + ) + + assert len(files) == 1 + try: + assert files[0].filename == expected_name + assert files[0].fp.name == str(video_path) + finally: + files[0].close() + + @pytest.mark.asyncio async def test_discord_group_message_includes_guild_and_channel_name(): adapter = DiscordPlatformAdapter.__new__(DiscordPlatformAdapter)