From 4a5f9ff4fb16f47979990bf36d5eaab0f86f3d6a Mon Sep 17 00:00:00 2001 From: SXP-Simon Date: Sun, 8 Feb 2026 16:47:01 +0800 Subject: [PATCH] feat: implement DiscordAdapter with py-cord integration --- .../platform/adapters/discord_adapter.py | 524 ++++++++++++++---- tests/unit_test_platform.py | 56 +- 2 files changed, 432 insertions(+), 148 deletions(-) diff --git a/src/infrastructure/platform/adapters/discord_adapter.py b/src/infrastructure/platform/adapters/discord_adapter.py index 5e32365..ae1c7ee 100644 --- a/src/infrastructure/platform/adapters/discord_adapter.py +++ b/src/infrastructure/platform/adapters/discord_adapter.py @@ -10,6 +10,13 @@ Discord 平台适配器 from datetime import datetime, timedelta from typing import List, Optional, Any, Dict +import asyncio +import logging + +try: + import discord +except ImportError: + discord = None from ....domain.value_objects.unified_message import ( UnifiedMessage, @@ -23,6 +30,7 @@ from ....domain.value_objects.platform_capabilities import ( from ....domain.value_objects.unified_group import UnifiedGroup, UnifiedMember from ..base import PlatformAdapter +logger = logging.getLogger(__name__) class DiscordAdapter(PlatformAdapter): """ @@ -42,6 +50,10 @@ class DiscordAdapter(PlatformAdapter): super().__init__(bot_instance, config) # 机器人自己的用户 ID,用于过滤消息 self.bot_user_id = str(config.get("bot_user_id", "")) if config else "" + + # 尝试从 bot 实例获取 ID + if not self.bot_user_id and hasattr(self.bot, "user") and self.bot.user: + self.bot_user_id = str(self.bot.user.id) def _init_capabilities(self) -> PlatformCapabilities: """初始化 Discord 平台能力""" @@ -68,113 +80,194 @@ class DiscordAdapter(PlatformAdapter): 返回: UnifiedMessage 列表 """ - # TODO: 实现 Discord 消息获取逻辑 - # 需要根据 AstrBot 的 Discord 集成方式来实现 - # 通常需要调用 Discord API 的 channel.history() 方法 - - # 示例实现框架: - # if not hasattr(self.bot, "get_channel"): - # return [] - # - # try: - # channel = self.bot.get_channel(int(group_id)) - # if not channel: - # return [] - # - # end_time = datetime.now() - # start_time = end_time - timedelta(days=days) - # - # messages = [] - # async for msg in channel.history(limit=max_count, after=start_time): - # if str(msg.author.id) == self.bot_user_id: - # continue - # unified = self._convert_message(msg, group_id) - # if unified: - # messages.append(unified) - # - # return messages - # except Exception: - # return [] - - return [] + if not discord: + logger.error("未安装 py-cord 库,无法使用 Discord 适配器") + return [] + + try: + channel_id = int(group_id) + channel = self.bot.get_channel(channel_id) + if not channel: + # 尝试 fetch (API调用) + try: + channel = await self.bot.fetch_channel(channel_id) + except Exception: + logger.warning(f"无法找到频道 ID: {group_id}") + return [] + + # 检查频道是否支持历史记录 + if not hasattr(channel, "history"): + logger.warning(f"频道 {group_id} 不支持历史消息获取") + return [] + + end_time = datetime.now() + start_time = end_time - timedelta(days=days) + + messages = [] + + # 构建 history 参数 + history_kwargs = { + "limit": max_count, + "after": start_time + } + if before_id: + try: + # before 可以接受 Message 对象或 ID (int) + history_kwargs["before"] = discord.Object(id=int(before_id)) + except ValueError: + pass + + # 获取消息 + async for msg in channel.history(**history_kwargs): + # 过滤机器人自己的消息(如果配置了 ID) + if self.bot_user_id and str(msg.author.id) == self.bot_user_id: + continue + + unified = self._convert_message(msg, group_id) + if unified: + messages.append(unified) + + # 按时间升序排序 + messages.sort(key=lambda m: m.timestamp) + return messages + + except Exception as e: + logger.error(f"获取 Discord 消息失败: {e}", exc_info=True) + return [] def _convert_message(self, raw_msg: Any, group_id: str) -> Optional[UnifiedMessage]: """ 将 Discord 消息转换为统一格式 参数: - raw_msg: Discord 原始消息对象 + raw_msg: Discord 原始消息对象 (discord.Message) group_id: 频道 ID 返回: UnifiedMessage 或 None """ - # TODO: 实现 Discord 消息转换逻辑 - # 示例: - # try: - # contents = [] - # - # # 文本内容 - # if raw_msg.content: - # contents.append(MessageContent( - # type=MessageContentType.TEXT, - # text=raw_msg.content - # )) - # - # # 图片附件 - # for attachment in raw_msg.attachments: - # if attachment.content_type and attachment.content_type.startswith("image/"): - # contents.append(MessageContent( - # type=MessageContentType.IMAGE, - # url=attachment.url - # )) - # - # return UnifiedMessage( - # message_id=str(raw_msg.id), - # sender_id=str(raw_msg.author.id), - # sender_name=raw_msg.author.display_name, - # sender_card=raw_msg.author.nick, - # group_id=group_id, - # text_content=raw_msg.content, - # contents=tuple(contents), - # timestamp=int(raw_msg.created_at.timestamp()), - # platform="discord", - # reply_to_id=str(raw_msg.reference.message_id) if raw_msg.reference else None, - # ) - # except Exception: - # return None - - return None + try: + contents = [] + + # 1. 文本内容 + if raw_msg.content: + contents.append(MessageContent( + type=MessageContentType.TEXT, + text=raw_msg.content + )) + + # 2. 附件处理 + for attachment in raw_msg.attachments: + content_type = attachment.content_type or "" + if content_type.startswith("image/"): + contents.append(MessageContent( + type=MessageContentType.IMAGE, + url=attachment.url + )) + elif content_type.startswith("video/"): + contents.append(MessageContent( + type=MessageContentType.VIDEO, + url=attachment.url + )) + elif content_type.startswith("audio/"): + contents.append(MessageContent( + type=MessageContentType.VOICE, + url=attachment.url + )) + else: + contents.append(MessageContent( + type=MessageContentType.FILE, + url=attachment.url, + raw_data={"filename": attachment.filename, "size": attachment.size} + )) + + # 3. 嵌入内容 (Embeds) - 通常是富文本或图片 + for embed in raw_msg.embeds: + if embed.image: + contents.append(MessageContent( + type=MessageContentType.IMAGE, + url=embed.image.url + )) + # 其他 embed 内容暂作为未知类型或文本处理 + if embed.description: + contents.append(MessageContent( + type=MessageContentType.TEXT, + text=f"\n[Embed] {embed.description}" + )) + + # 4. 贴纸 (Stickers) + if raw_msg.stickers: + for sticker in raw_msg.stickers: + contents.append(MessageContent( + type=MessageContentType.IMAGE, # 贴纸视为图片 + url=sticker.url, + raw_data={"sticker_id": str(sticker.id), "sticker_name": sticker.name} + )) + + # 发送者名片 (昵称) + sender_card = None + if hasattr(raw_msg.author, "nick") and raw_msg.author.nick: + sender_card = raw_msg.author.nick + elif hasattr(raw_msg.author, "global_name") and raw_msg.author.global_name: + sender_card = raw_msg.author.global_name + + return UnifiedMessage( + message_id=str(raw_msg.id), + sender_id=str(raw_msg.author.id), + sender_name=raw_msg.author.name, # 用户名 + sender_card=sender_card, # 服务器昵称 + group_id=group_id, + text_content=raw_msg.content, + contents=tuple(contents), + timestamp=int(raw_msg.created_at.timestamp()), + platform="discord", + reply_to_id=str(raw_msg.reference.message_id) if raw_msg.reference else None, + ) + except Exception as e: + logger.error(f"转换 Discord 消息失败: {e}") + return None def convert_to_raw_format(self, messages: List[UnifiedMessage]) -> List[dict]: """ - 将统一消息格式转换为 Discord 原生格式 + 将统一消息格式转换为 Discord 原生格式 (模拟) 用于与现有分析器的向后兼容。 - Discord 格式与 OneBot 不同,这里返回通用字典格式。 """ raw_messages = [] for msg in messages: - # Discord 风格的消息格式 + # 构造模拟的 Discord 消息字典 raw_msg = { "id": msg.message_id, "channel_id": msg.group_id, "author": { "id": msg.sender_id, "username": msg.sender_name, - "nick": msg.sender_card, + "discriminator": "0000", # 兼容旧格式 + "global_name": msg.sender_card, + "avatar": None, # 暂不获取头像hash }, "content": msg.text_content, - "timestamp": msg.timestamp, + "timestamp": datetime.fromtimestamp(msg.timestamp).isoformat(), + "edited_timestamp": None, + "tts": False, + "mention_everyone": False, + "mentions": [], + "mention_roles": [], "attachments": [], "embeds": [], + "pinned": False, + "type": 0, } # 处理附件 for content in msg.contents: if content.type == MessageContentType.IMAGE: raw_msg["attachments"].append({ + "id": "0", # 伪造ID + "filename": "image.png", + "size": 0, "url": content.url, + "proxy_url": content.url, "content_type": "image/png", }) @@ -191,23 +284,33 @@ class DiscordAdapter(PlatformAdapter): reply_to: Optional[str] = None, ) -> bool: """发送文本消息到 Discord 频道""" - # TODO: 实现 Discord 消息发送 - # 示例: - # try: - # channel = self.bot.get_channel(int(group_id)) - # if not channel: - # return False - # - # if reply_to: - # ref_msg = await channel.fetch_message(int(reply_to)) - # await channel.send(text, reference=ref_msg) - # else: - # await channel.send(text) - # return True - # except Exception: - # return False - - return False + if not discord: return False + + try: + channel_id = int(group_id) + channel = self.bot.get_channel(channel_id) + if not channel: + channel = await self.bot.fetch_channel(channel_id) + + if not hasattr(channel, "send"): + return False + + reference = None + if reply_to: + try: + # 创建 MessageReference + reference = discord.MessageReference( + message_id=int(reply_to), + channel_id=channel_id + ) + except ValueError: + pass + + await channel.send(content=text, reference=reference) + return True + except Exception as e: + logger.error(f"Discord 发送文本失败: {e}") + return False async def send_image( self, @@ -216,8 +319,34 @@ class DiscordAdapter(PlatformAdapter): caption: str = "", ) -> bool: """发送图片到 Discord 频道""" - # TODO: 实现 Discord 图片发送 - return False + if not discord: return False + + try: + channel_id = int(group_id) + channel = self.bot.get_channel(channel_id) + if not channel: + channel = await self.bot.fetch_channel(channel_id) + + if not hasattr(channel, "send"): + return False + + # 处理本地文件或 URL + file_to_send = None + if image_path.startswith(("http://", "https://")): + # URL 方式,直接放在内容里或者作为 embed (Discord.py send 不直接支持 url 作为 file) + # 简单起见,如果是有 caption,将 URL 拼接到 content + content = f"{caption}\n{image_path}" if caption else image_path + await channel.send(content=content) + return True + else: + # 本地文件 + file_to_send = discord.File(image_path) + await channel.send(content=caption, file=file_to_send) + return True + + except Exception as e: + logger.error(f"Discord 发送图片失败: {e}") + return False async def send_file( self, @@ -226,41 +355,122 @@ class DiscordAdapter(PlatformAdapter): filename: Optional[str] = None, ) -> bool: """发送文件到 Discord 频道""" - # TODO: 实现 Discord 文件发送 - return False + if not discord: return False + + try: + channel_id = int(group_id) + channel = self.bot.get_channel(channel_id) + if not channel: + channel = await self.bot.fetch_channel(channel_id) + + if not hasattr(channel, "send"): + return False + + file_to_send = discord.File(file_path, filename=filename) + await channel.send(file=file_to_send) + return True + except Exception as e: + logger.error(f"Discord 发送文件失败: {e}") + return False # ==================== IGroupInfoRepository ==================== async def get_group_info(self, group_id: str) -> Optional[UnifiedGroup]: """获取 Discord 频道信息""" - # TODO: 实现 Discord 频道信息获取 - # 示例: - # try: - # channel = self.bot.get_channel(int(group_id)) - # if not channel: - # return None - # - # return UnifiedGroup( - # group_id=str(channel.id), - # group_name=channel.name, - # member_count=channel.guild.member_count if hasattr(channel, "guild") else 0, - # owner_id=str(channel.guild.owner_id) if hasattr(channel, "guild") else None, - # platform="discord", - # ) - # except Exception: - # return None + if not discord: return None - return None + try: + channel_id = int(group_id) + channel = self.bot.get_channel(channel_id) + if not channel: + channel = await self.bot.fetch_channel(channel_id) + + # 尝试获取 Guild 信息 + guild = getattr(channel, "guild", None) + + group_name = getattr(channel, "name", str(channel.id)) + if guild: + # 如果是公会频道,可以用 Guild 信息补充 + member_count = guild.member_count + owner_id = str(guild.owner_id) + else: + # 私信或群组私信 + member_count = len(getattr(channel, "recipients", [])) + 1 # +1 for bot + owner_id = str(getattr(channel, "owner_id", "")) + + return UnifiedGroup( + group_id=str(channel.id), + group_name=group_name, + member_count=member_count, + owner_id=owner_id or None, + create_time=int(channel.created_at.timestamp()), + platform="discord", + ) + except Exception as e: + logger.error(f"Discord 获取群组信息失败: {e}") + return None async def get_group_list(self) -> List[str]: - """获取机器人所在的所有频道 ID""" - # TODO: 实现 Discord 频道列表获取 - return [] + """获取机器人所在的所有频道 ID (仅列出 TextChannel)""" + if not discord: return [] + + try: + # 遍历所有 Guilds 和 Channels + channel_ids = [] + for guild in self.bot.guilds: + for channel in guild.text_channels: + channel_ids.append(str(channel.id)) + return channel_ids + except Exception as e: + logger.error(f"Discord 获取群组列表失败: {e}") + return [] async def get_member_list(self, group_id: str) -> List[UnifiedMember]: """获取 Discord 服务器成员列表""" - # TODO: 实现 Discord 成员列表获取 - return [] + if not discord: return [] + + try: + channel_id = int(group_id) + channel = self.bot.get_channel(channel_id) + if not channel: + channel = await self.bot.fetch_channel(channel_id) + + guild = getattr(channel, "guild", None) + if not guild: + # 非公会频道(如 DM),返回收件人 + members = [] + for user in getattr(channel, "recipients", []): + members.append(UnifiedMember( + user_id=str(user.id), + nickname=user.display_name, + card=None, + role="member", + join_time=None + )) + return members + + # 公会频道 + members = [] + # 注意:如果 member_count 很大,members 可能不全(取决于 intent 和 cache) + # 需要启用 GUILD_MEMBERS intent + for member in guild.members: + role = "member" + if member.id == guild.owner_id: + role = "owner" + elif member.guild_permissions.administrator: + role = "admin" + + members.append(UnifiedMember( + user_id=str(member.id), + nickname=member.name, + card=member.nick or member.global_name, # 优先显示服务器昵称 + role=role, + join_time=int(member.joined_at.timestamp()) if member.joined_at else None, + )) + return members + except Exception as e: + logger.error(f"Discord 获取成员列表失败: {e}") + return [] async def get_member_info( self, @@ -268,8 +478,49 @@ class DiscordAdapter(PlatformAdapter): user_id: str, ) -> Optional[UnifiedMember]: """获取特定成员信息""" - # TODO: 实现 Discord 成员信息获取 - return None + if not discord: return None + + try: + channel_id = int(group_id) + channel = self.bot.get_channel(channel_id) + if not channel: + channel = await self.bot.fetch_channel(channel_id) + + guild = getattr(channel, "guild", None) + if not guild: + # 私信,尝试 fetch user + user = await self.bot.fetch_user(int(user_id)) + return UnifiedMember( + user_id=str(user.id), + nickname=user.name, + card=user.display_name, + role="member", + join_time=None + ) + + member = guild.get_member(int(user_id)) + if not member: + member = await guild.fetch_member(int(user_id)) + + if not member: + return None + + role = "member" + if member.id == guild.owner_id: + role = "owner" + elif member.guild_permissions.administrator: + role = "admin" + + return UnifiedMember( + user_id=str(member.id), + nickname=member.name, + card=member.nick or member.global_name, + role=role, + join_time=int(member.joined_at.timestamp()) if member.joined_at else None, + ) + except Exception as e: + logger.error(f"Discord 获取成员信息失败: {e}") + return None # ==================== IAvatarRepository ==================== @@ -279,10 +530,23 @@ class DiscordAdapter(PlatformAdapter): size: int = 100, ) -> Optional[str]: """获取 Discord 用户头像 URL""" - # Discord 头像 URL 格式 - # https://cdn.discordapp.com/avatars/{user_id}/{avatar_hash}.png?size={size} - # 需要知道用户的 avatar_hash,这里返回默认头像 - return f"https://cdn.discordapp.com/embed/avatars/{int(user_id) % 5}.png" + if not discord: return None + + try: + user = self.bot.get_user(int(user_id)) + if not user: + user = await self.bot.fetch_user(int(user_id)) + + if user: + # 调整 size 到最接近的 2 的幂次方 + allowed_sizes = [16, 32, 64, 128, 256, 512, 1024, 2048, 4096] + target_size = min(allowed_sizes, key=lambda x: abs(x - size)) + + # display_avatar 自动处理默认头像 + return user.display_avatar.with_size(target_size).url + return None + except Exception: + return None async def get_user_avatar_data( self, @@ -290,7 +554,7 @@ class DiscordAdapter(PlatformAdapter): size: int = 100, ) -> Optional[str]: """获取 Discord 用户头像 Base64 数据""" - # TODO: 实现头像数据获取 + # 暂时只返回 None,让上层使用 URL return None async def get_group_avatar_url( @@ -299,8 +563,22 @@ class DiscordAdapter(PlatformAdapter): size: int = 100, ) -> Optional[str]: """获取 Discord 服务器图标 URL""" - # TODO: 实现服务器图标获取 - return None + if not discord: return None + + try: + channel_id = int(group_id) + channel = self.bot.get_channel(channel_id) + if not channel: + channel = await self.bot.fetch_channel(channel_id) + + guild = getattr(channel, "guild", None) + if guild and guild.icon: + allowed_sizes = [16, 32, 64, 128, 256, 512, 1024, 2048, 4096] + target_size = min(allowed_sizes, key=lambda x: abs(x - size)) + return guild.icon.with_size(target_size).url + return None + except Exception: + return None async def batch_get_avatar_urls( self, diff --git a/tests/unit_test_platform.py b/tests/unit_test_platform.py index 61f1e79..9b51354 100644 --- a/tests/unit_test_platform.py +++ b/tests/unit_test_platform.py @@ -60,42 +60,48 @@ class TestPlatformArchitecture(unittest.TestCase): adapter = PlatformAdapterFactory.create("discord", MagicMock(), {}) self.assertIsInstance(adapter, DiscordAdapter) - def test_orchestrator_raw_conversion(self): - """测试编排器的原始格式转换 (验证硬编码移除)""" - # Mock adapter - mock_adapter = MagicMock(spec=PlatformAdapter) - mock_adapter.get_capabilities.return_value = DISCORD_CAPABILITIES + def test_discord_fetch_messages(self): + """测试 Discord 消息获取逻辑 (Mocked)""" + # Mock bot instance + mock_bot = MagicMock() + mock_channel = MagicMock() + mock_bot.get_channel.return_value = mock_channel - # Mock fetch_messages return - mock_msg = UnifiedMessage( - message_id="1", sender_id="u1", sender_name="User", - group_id="g1", text_content="test", contents=[], - timestamp=1234567890, platform="discord" - ) - mock_adapter.fetch_messages = AsyncMock(return_value=[mock_msg]) + # Mock message history + # Create a mock message that mimics discord.Message + mock_msg = MagicMock() + mock_msg.id = 12345 + mock_msg.content = "test message" + mock_msg.author.id = 999 + mock_msg.author.name = "User" + mock_msg.created_at.timestamp.return_value = 1600000000 + mock_msg.attachments = [] + mock_msg.embeds = [] + mock_msg.stickers = [] + mock_msg.reference = None - # Mock convert_to_raw_format - mock_adapter.convert_to_raw_format.return_value = [{"id": "1", "content": "raw"}] + # history returns an async iterator + async def async_iter(): + yield mock_msg + + mock_channel.history.return_value = async_iter() - # Create orchestrator - orchestrator = AnalysisOrchestrator(adapter=mock_adapter) + # Initialize adapter + adapter = DiscordAdapter(bot_instance=mock_bot, config={"bot_user_id": "123"}) - # Run sync wrapper for async method (simplified for unit test structure) + # Run async test import asyncio loop = asyncio.new_event_loop() asyncio.set_event_loop(loop) - # Call fetch_messages_as_raw - raw_msgs = loop.run_until_complete( - orchestrator.fetch_messages_as_raw("g1") + messages = loop.run_until_complete( + adapter.fetch_messages("1001", days=1) ) - # Verify result - self.assertEqual(len(raw_msgs), 1) - self.assertEqual(raw_msgs[0]["content"], "raw") + self.assertEqual(len(messages), 1) + self.assertEqual(messages[0].text_content, "test message") + self.assertEqual(messages[0].platform, "discord") - # Verify adapter method was called (PROVING adapter pattern is used) - mock_adapter.convert_to_raw_format.assert_called_once() loop.close() if __name__ == "__main__":