feat: implement DiscordAdapter with py-cord integration

This commit is contained in:
SXP-Simon
2026-02-08 16:47:01 +08:00
parent a2adfe6561
commit 4a5f9ff4fb
2 changed files with 432 additions and 148 deletions
@@ -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,
+31 -25
View File
@@ -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__":