mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-05 15:38:19 +08:00
支持 Slack 和 Discord 自动注册命令
This commit is contained in:
@@ -1,13 +1,22 @@
|
||||
import copy
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from urllib.parse import quote, unquote
|
||||
from typing import Optional, Union, List, Tuple, Any
|
||||
|
||||
from app.core.context import MediaInfo, Context
|
||||
from app.core.event import eventmanager
|
||||
from app.log import logger
|
||||
from app.modules import _ModuleBase, _MessageBase
|
||||
from app.schemas import MessageChannel, CommingMessage, Notification, MessageResponse
|
||||
from app.schemas.types import ModuleType
|
||||
from app.schemas import (
|
||||
CommandRegisterEventData,
|
||||
CommingMessage,
|
||||
MessageChannel,
|
||||
MessageResponse,
|
||||
Notification,
|
||||
)
|
||||
from app.schemas.types import ChainEventType, ModuleType
|
||||
from app.utils.http import RequestUtils
|
||||
from app.utils.structures import DictUtils
|
||||
|
||||
try:
|
||||
from app.modules.discord.discord import Discord
|
||||
@@ -530,6 +539,54 @@ class DiscordModule(_ModuleBase, _MessageBase[Discord]):
|
||||
return True
|
||||
return False
|
||||
|
||||
def register_commands(self, commands: Dict[str, dict]) -> None:
|
||||
"""
|
||||
注册命令,实现这个函数接收系统可用的命令菜单。
|
||||
|
||||
:param commands: 命令字典
|
||||
"""
|
||||
for client_config in self.get_configs().values():
|
||||
client = self.get_instance(client_config.name)
|
||||
if not client:
|
||||
continue
|
||||
|
||||
scoped_commands = copy.deepcopy(commands)
|
||||
event = eventmanager.send_event(
|
||||
ChainEventType.CommandRegister,
|
||||
CommandRegisterEventData(
|
||||
commands=scoped_commands,
|
||||
origin="Discord",
|
||||
service=client_config.name,
|
||||
),
|
||||
)
|
||||
|
||||
if event and event.event_data:
|
||||
event_data: CommandRegisterEventData = event.event_data
|
||||
if event_data.cancel:
|
||||
client.delete_commands()
|
||||
logger.debug(
|
||||
f"Command registration for {client_config.name} canceled by event: {event_data.source}"
|
||||
)
|
||||
continue
|
||||
scoped_commands = event_data.commands or {}
|
||||
if not scoped_commands:
|
||||
logger.debug("Filtered commands are empty, skipping registration.")
|
||||
client.delete_commands()
|
||||
|
||||
filtered_scoped_commands = DictUtils.filter_keys_to_subset(
|
||||
scoped_commands,
|
||||
commands,
|
||||
)
|
||||
if not filtered_scoped_commands:
|
||||
logger.debug("Filtered commands are empty, skipping registration.")
|
||||
client.delete_commands()
|
||||
continue
|
||||
if filtered_scoped_commands != commands:
|
||||
logger.debug(
|
||||
f"Command set has changed, Updating new commands: {filtered_scoped_commands}"
|
||||
)
|
||||
client.register_commands(filtered_scoped_commands)
|
||||
|
||||
def mark_message_processing_started(
|
||||
self,
|
||||
channel: MessageChannel,
|
||||
|
||||
@@ -31,6 +31,8 @@ class Discord:
|
||||
Discord Bot 通知与交互实现(基于 discord.py 2.6.4)
|
||||
"""
|
||||
|
||||
_MAX_SLASH_COMMANDS = 100
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
DISCORD_BOT_TOKEN: Optional[str] = None,
|
||||
@@ -69,7 +71,7 @@ class Discord:
|
||||
self._client: Optional[discord.Client] = discord.Client(
|
||||
intents=intents, proxy=settings.PROXY_HOST
|
||||
)
|
||||
self._tree: Optional[app_commands.CommandTree] = None
|
||||
self._tree: Optional[app_commands.CommandTree] = app_commands.CommandTree(self._client)
|
||||
self._loop: asyncio.AbstractEventLoop = asyncio.new_event_loop()
|
||||
self._thread: Optional[threading.Thread] = None
|
||||
self._ready_event = threading.Event()
|
||||
@@ -84,6 +86,7 @@ class Discord:
|
||||
self._typing_interval_seconds = 5
|
||||
self._typing_initial_delay_seconds = 1
|
||||
self._typing_max_duration_seconds = 10 * 60
|
||||
self._registered_commands: Optional[Dict[str, dict]] = None
|
||||
|
||||
self._register_events()
|
||||
self._start()
|
||||
@@ -101,6 +104,11 @@ class Discord:
|
||||
self._bot_user_id = self._client.user.id if self._client.user else None
|
||||
self._ready_event.set()
|
||||
logger.info(f"Discord Bot 已登录:{self._client.user}")
|
||||
if self._registered_commands is not None:
|
||||
try:
|
||||
await self._sync_registered_commands()
|
||||
except Exception as err:
|
||||
logger.error(f"同步 Discord 斜杠命令失败:{err}")
|
||||
|
||||
@self._client.event
|
||||
async def on_message(message: discord.Message):
|
||||
@@ -232,6 +240,169 @@ class Discord:
|
||||
def get_state(self) -> bool:
|
||||
return self._ready_event.is_set() and self._client is not None
|
||||
|
||||
def register_commands(self, commands: Dict[str, dict]) -> bool:
|
||||
"""
|
||||
注册 Discord 斜杠命令。
|
||||
|
||||
:param commands: 命令字典,键为斜杠命令,值包含描述和分类等元数据
|
||||
:return: 是否成功提交同步任务
|
||||
"""
|
||||
self._registered_commands = dict(commands or {})
|
||||
return self._schedule_command_sync()
|
||||
|
||||
def delete_commands(self) -> bool:
|
||||
"""
|
||||
清理 Discord 斜杠命令。
|
||||
|
||||
:return: 是否成功提交同步任务
|
||||
"""
|
||||
self._registered_commands = {}
|
||||
return self._schedule_command_sync()
|
||||
|
||||
def _schedule_command_sync(self) -> bool:
|
||||
"""在 Discord 事件循环中提交命令同步任务。"""
|
||||
if not self._tree or not self._loop:
|
||||
return False
|
||||
if not self.get_state():
|
||||
logger.debug("Discord Bot 未就绪,斜杠命令将在登录后同步")
|
||||
return True
|
||||
try:
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
self._sync_registered_commands(), self._loop
|
||||
)
|
||||
return bool(future.result(timeout=30))
|
||||
except Exception as err:
|
||||
logger.error(f"同步 Discord 斜杠命令失败:{err}")
|
||||
return False
|
||||
|
||||
async def _sync_registered_commands(self) -> bool:
|
||||
"""将当前命令集合同步到 Discord 应用命令树。"""
|
||||
if not self._tree or not self._client:
|
||||
return False
|
||||
if not self._client.is_ready():
|
||||
await self._client.wait_until_ready()
|
||||
|
||||
guild = discord.Object(id=self._guild_id) if self._guild_id else None
|
||||
self._tree.clear_commands(guild=guild)
|
||||
|
||||
commands = self._registered_commands or {}
|
||||
registered_count = 0
|
||||
seen_names = set()
|
||||
for command_text, command_data in commands.items():
|
||||
if registered_count >= self._MAX_SLASH_COMMANDS:
|
||||
logger.warning(
|
||||
f"Discord 斜杠命令数量超过 {self._MAX_SLASH_COMMANDS} 个,后续命令已跳过"
|
||||
)
|
||||
break
|
||||
command_name = self._normalize_slash_command_name(command_text)
|
||||
if not command_name or command_name in seen_names:
|
||||
logger.warning(f"跳过无效或重复的 Discord 斜杠命令:{command_text}")
|
||||
continue
|
||||
seen_names.add(command_name)
|
||||
description = self._normalize_slash_command_description(
|
||||
command_data.get("description") if isinstance(command_data, dict) else None,
|
||||
command_name,
|
||||
)
|
||||
self._tree.add_command(
|
||||
self._build_slash_command(command_text, command_name, description),
|
||||
guild=guild,
|
||||
override=True,
|
||||
)
|
||||
registered_count += 1
|
||||
|
||||
synced_commands = await self._tree.sync(guild=guild)
|
||||
logger.info(f"Discord 斜杠命令已同步:{len(synced_commands)} 个")
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _normalize_slash_command_name(command_text: str) -> str:
|
||||
"""转换为 Discord 允许的斜杠命令名称。"""
|
||||
command_name = str(command_text or "").strip().lstrip("/").lower()
|
||||
if not re.fullmatch(r"[a-z0-9_-]{1,32}", command_name):
|
||||
return ""
|
||||
return command_name
|
||||
|
||||
@staticmethod
|
||||
def _normalize_slash_command_description(
|
||||
description: Optional[str],
|
||||
fallback: str,
|
||||
) -> str:
|
||||
"""整理 Discord 斜杠命令描述,满足长度要求。"""
|
||||
normalized = str(description or fallback or "MoviePilot").strip()
|
||||
return normalized[:100] or "MoviePilot"
|
||||
|
||||
def _build_slash_command(
|
||||
self,
|
||||
command_text: str,
|
||||
command_name: str,
|
||||
description: str,
|
||||
) -> app_commands.Command:
|
||||
"""构建 Discord 斜杠命令对象。"""
|
||||
|
||||
async def _callback(
|
||||
interaction: discord.Interaction,
|
||||
args: Optional[str] = None,
|
||||
) -> None:
|
||||
await self._handle_slash_command(interaction, command_text, args)
|
||||
|
||||
_callback.__name__ = f"moviepilot_{command_name}"
|
||||
_callback = app_commands.describe(args="命令参数")(_callback)
|
||||
return app_commands.Command(
|
||||
name=command_name,
|
||||
description=description,
|
||||
callback=_callback,
|
||||
)
|
||||
|
||||
async def _handle_slash_command(
|
||||
self,
|
||||
interaction: discord.Interaction,
|
||||
command_text: str,
|
||||
args: Optional[str] = None,
|
||||
) -> None:
|
||||
"""处理 Discord 斜杠命令回调,并转发到统一消息入口。"""
|
||||
try:
|
||||
await interaction.response.defer(ephemeral=True, thinking=True)
|
||||
except Exception as err:
|
||||
logger.debug(f"延迟响应 Discord 斜杠命令失败:{err}")
|
||||
|
||||
userid = str(interaction.user.id) if interaction.user else None
|
||||
chat_id = str(interaction.channel.id) if interaction.channel else None
|
||||
username = None
|
||||
if interaction.user:
|
||||
username = (
|
||||
getattr(interaction.user, "display_name", None)
|
||||
or getattr(interaction.user, "global_name", None)
|
||||
or getattr(interaction.user, "name", None)
|
||||
)
|
||||
if userid and chat_id:
|
||||
self._update_user_chat_mapping(userid, chat_id)
|
||||
|
||||
arg_text = str(args or "").strip()
|
||||
payload = {
|
||||
"type": "message",
|
||||
"userid": userid,
|
||||
"username": username,
|
||||
"user_tag": str(interaction.user) if interaction.user else None,
|
||||
"text": f"{command_text} {arg_text}".strip(),
|
||||
"message_id": str(interaction.id),
|
||||
"chat_id": chat_id,
|
||||
"channel_type": "dm"
|
||||
if isinstance(interaction.channel, discord.DMChannel)
|
||||
else "guild",
|
||||
}
|
||||
await self._post_to_ds(payload)
|
||||
|
||||
try:
|
||||
if interaction.response.is_done():
|
||||
await interaction.followup.send("命令已提交,请稍等...", ephemeral=True)
|
||||
else:
|
||||
await interaction.response.send_message(
|
||||
"命令已提交,请稍等...",
|
||||
ephemeral=True,
|
||||
)
|
||||
except Exception as err:
|
||||
logger.debug(f"发送 Discord 斜杠命令确认失败:{err}")
|
||||
|
||||
def send_msg(
|
||||
self,
|
||||
title: str,
|
||||
|
||||
Reference in New Issue
Block a user