diff --git a/src/infrastructure/platform/adapters/discord_adapter.py b/src/infrastructure/platform/adapters/discord_adapter.py index 7f1d2d5..3fb3fdb 100644 --- a/src/infrastructure/platform/adapters/discord_adapter.py +++ b/src/infrastructure/platform/adapters/discord_adapter.py @@ -32,16 +32,17 @@ from ..base import PlatformAdapter logger = logging.getLogger(__name__) + class DiscordAdapter(PlatformAdapter): """ Discord 平台适配器 - + 实现 PlatformAdapter 接口,提供 Discord 平台的消息操作。 - + 使用方式: 1. 通过 PlatformAdapterFactory.create("discord", bot_instance, config) 创建 2. 或直接实例化:DiscordAdapter(bot_instance, config) - + 配置参数: - bot_user_id: 机器人的 Discord 用户 ID(用于过滤自己的消息) """ @@ -50,10 +51,35 @@ 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) + + # 获取实际的 Discord 客户端 + self._discord_client = self._get_discord_client() + + # 尝试从 Discord 客户端获取 ID + if not self.bot_user_id and self._discord_client: + if hasattr(self._discord_client, "user") and self._discord_client.user: + self.bot_user_id = str(self._discord_client.user.id) + + def _get_discord_client(self) -> Any: + """ + 获取实际的 Discord 客户端实例 + + AstrBot 的 DiscordPlatformAdapter 将 Discord 客户端存储在 self.client 中 + """ + # 如果 bot 本身就是 Discord client (有 get_channel 方法) + if hasattr(self.bot, "get_channel"): + return self.bot + # 如果 bot 是 DiscordPlatformAdapter,client 在 self.bot.client 中 + if hasattr(self.bot, "client"): + return self.bot.client + # 尝试其他可能的属性名 + for attr in ["_client", "discord_client", "_discord_client"]: + if hasattr(self.bot, attr): + client = getattr(self.bot, attr) + if hasattr(client, "get_channel"): + return client + logger.warning(f"无法从 {type(self.bot).__name__} 获取 Discord 客户端") + return None def _init_capabilities(self) -> PlatformCapabilities: """初始化 Discord 平台能力""" @@ -70,13 +96,13 @@ class DiscordAdapter(PlatformAdapter): ) -> List[UnifiedMessage]: """ 获取 Discord 频道消息历史 - + 参数: group_id: Discord 频道 ID days: 获取多少天内的消息 max_count: 最大消息数量 before_id: 从此消息 ID 之前开始获取(用于分页) - + 返回: UnifiedMessage 列表 """ @@ -86,30 +112,27 @@ class DiscordAdapter(PlatformAdapter): try: channel_id = int(group_id) - channel = self.bot.get_channel(channel_id) + channel = self._discord_client.get_channel(channel_id) if not channel: # 尝试 fetch (API调用) try: - channel = await self.bot.fetch_channel(channel_id) + channel = await self._discord_client.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 - } + history_kwargs = {"limit": max_count, "after": start_time} if before_id: try: # before 可以接受 Message 对象或 ID (int) @@ -122,15 +145,15 @@ class DiscordAdapter(PlatformAdapter): # 过滤机器人自己的消息(如果配置了 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 [] @@ -138,72 +161,87 @@ class DiscordAdapter(PlatformAdapter): def _convert_message(self, raw_msg: Any, group_id: str) -> Optional[UnifiedMessage]: """ 将 Discord 消息转换为统一格式 - + 参数: raw_msg: Discord 原始消息对象 (discord.Message) group_id: 频道 ID - + 返回: UnifiedMessage 或 None """ try: contents = [] - + # 1. 文本内容 if raw_msg.content: - contents.append(MessageContent( - type=MessageContentType.TEXT, - text=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 - )) + contents.append( + MessageContent( + type=MessageContentType.IMAGE, url=attachment.url + ) + ) elif content_type.startswith("video/"): - contents.append(MessageContent( - type=MessageContentType.VIDEO, - url=attachment.url - )) + contents.append( + MessageContent( + type=MessageContentType.VIDEO, url=attachment.url + ) + ) elif content_type.startswith("audio/"): - contents.append(MessageContent( - type=MessageContentType.VOICE, - url=attachment.url - )) + 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} - )) - + 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 - )) + 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}" - )) + 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} - )) - + 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: @@ -214,14 +252,16 @@ class DiscordAdapter(PlatformAdapter): 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, # 服务器昵称 + 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, + reply_to_id=str(raw_msg.reference.message_id) + if raw_msg.reference + else None, ) except Exception as e: logger.error(f"转换 Discord 消息失败: {e}") @@ -245,22 +285,36 @@ class DiscordAdapter(PlatformAdapter): }, "message": [], } - + # 构造消息链 for content in msg.contents: if content.type == MessageContentType.TEXT: - raw_msg["message"].append({"type": "text", "data": {"text": content.text}}) + raw_msg["message"].append( + {"type": "text", "data": {"text": content.text}} + ) elif content.type == MessageContentType.IMAGE: - raw_msg["message"].append({"type": "image", "data": {"url": content.url, "file": content.url}}) + raw_msg["message"].append( + { + "type": "image", + "data": {"url": content.url, "file": content.url}, + } + ) elif content.type == MessageContentType.AT: - raw_msg["message"].append({"type": "at", "data": {"qq": content.at_user_id}}) + raw_msg["message"].append( + {"type": "at", "data": {"qq": content.at_user_id}} + ) elif content.type == MessageContentType.REPLY: if content.raw_data and "reply_id" in content.raw_data: - raw_msg["message"].append({"type": "reply", "data": {"id": content.raw_data["reply_id"]}}) + raw_msg["message"].append( + { + "type": "reply", + "data": {"id": content.raw_data["reply_id"]}, + } + ) # 其他类型暂忽略或作为未知 - + raw_messages.append(raw_msg) - + return raw_messages # ==================== IMessageSender ==================== @@ -272,28 +326,28 @@ class DiscordAdapter(PlatformAdapter): reply_to: Optional[str] = None, ) -> bool: """发送文本消息到 Discord 频道""" - if not 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 - + reference = None if reply_to: try: # 创建 MessageReference reference = discord.MessageReference( - message_id=int(reply_to), - channel_id=channel_id + 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: @@ -307,14 +361,15 @@ class DiscordAdapter(PlatformAdapter): caption: str = "", ) -> bool: """发送图片到 Discord 频道""" - if not 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 @@ -331,7 +386,7 @@ class DiscordAdapter(PlatformAdapter): 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 @@ -343,8 +398,9 @@ class DiscordAdapter(PlatformAdapter): filename: Optional[str] = None, ) -> bool: """发送文件到 Discord 频道""" - if not discord: return False - + if not discord: + return False + try: channel_id = int(group_id) channel = self.bot.get_channel(channel_id) @@ -353,7 +409,7 @@ class DiscordAdapter(PlatformAdapter): 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 @@ -365,17 +421,18 @@ class DiscordAdapter(PlatformAdapter): async def get_group_info(self, group_id: str) -> Optional[UnifiedGroup]: """获取 Discord 频道信息""" - if not 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 信息 guild = getattr(channel, "guild", None) - + group_name = getattr(channel, "name", str(channel.id)) if guild: # 如果是公会频道,可以用 Guild 信息补充 @@ -383,9 +440,9 @@ class DiscordAdapter(PlatformAdapter): owner_id = str(guild.owner_id) else: # 私信或群组私信 - member_count = len(getattr(channel, "recipients", [])) + 1 # +1 for bot + 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, @@ -400,14 +457,15 @@ class DiscordAdapter(PlatformAdapter): async def get_group_list(self) -> List[str]: """获取机器人所在的所有频道 ID (仅列出 TextChannel)""" - if not discord: return [] - + if not discord: + return [] + try: # 遍历所有 Guilds 和 Channels channel_ids = [] - for guild in self.bot.guilds: + for guild in self._discord_client.guilds: for channel in guild.text_channels: - channel_ids.append(str(channel.id)) + channel_ids.append(str(channel.id)) return channel_ids except Exception as e: logger.error(f"Discord 获取群组列表失败: {e}") @@ -415,26 +473,29 @@ class DiscordAdapter(PlatformAdapter): async def get_member_list(self, group_id: str) -> List[UnifiedMember]: """获取 Discord 服务器成员列表""" - if not 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 - )) + members.append( + UnifiedMember( + user_id=str(user.id), + nickname=user.display_name, + card=None, + role="member", + join_time=None, + ) + ) return members # 公会频道 @@ -447,14 +508,18 @@ class DiscordAdapter(PlatformAdapter): 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, - )) + + 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}") @@ -466,14 +531,15 @@ class DiscordAdapter(PlatformAdapter): user_id: str, ) -> Optional[UnifiedMember]: """获取特定成员信息""" - if not 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 @@ -483,28 +549,30 @@ class DiscordAdapter(PlatformAdapter): nickname=user.name, card=user.display_name, role="member", - join_time=None + 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, + join_time=int(member.joined_at.timestamp()) + if member.joined_at + else None, ) except Exception as e: logger.error(f"Discord 获取成员信息失败: {e}") @@ -518,18 +586,19 @@ class DiscordAdapter(PlatformAdapter): size: int = 100, ) -> Optional[str]: """获取 Discord 用户头像 URL""" - if not discord: return None - + if not discord: + return None + try: - user = self.bot.get_user(int(user_id)) + user = self._discord_client.get_user(int(user_id)) if not user: - user = await self.bot.fetch_user(int(user_id)) - + user = await self._discord_client.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 @@ -551,19 +620,20 @@ class DiscordAdapter(PlatformAdapter): size: int = 100, ) -> Optional[str]: """获取 Discord 服务器图标 URL""" - if not 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 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 + 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