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
18 changes: 18 additions & 0 deletions astrbot/core/platform/sources/discord/discord_platform_event.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -16,6 +17,7 @@
Plain,
Record,
Reply,
Video,
)
from astrbot.api.platform import (
AstrBotMessage,
Expand Down Expand Up @@ -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()
Expand Down
89 changes: 86 additions & 3 deletions tests/test_discord_adapter.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand All @@ -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)
Expand Down
Loading