refactor(platforms): improve type safety in platform adapters

- Use ComponentType enum instead of string literals for component types
- Add type annotations for discord button declarations
- Clean up unnecessary code in various platform adapters
This commit is contained in:
LIghtJUNction
2026-03-31 20:15:46 +08:00
parent 3eccbca2c9
commit 65841684de
16 changed files with 49 additions and 49 deletions
@@ -1,13 +1,13 @@
import discord
from astrbot.api.message_components import BaseMessageComponent
from astrbot.api.message_components import BaseMessageComponent, ComponentType
# Discord专用组件
class DiscordEmbed(BaseMessageComponent):
"""Discord Embed消息组件"""
type: str = "discord_embed"
type: ComponentType = ComponentType.DiscordEmbed
def __init__(
self,
@@ -61,7 +61,7 @@ class DiscordEmbed(BaseMessageComponent):
class DiscordButton(BaseMessageComponent):
"""Discord按钮组件"""
type: str = "discord_button"
type: ComponentType = ComponentType.DiscordButton
def __init__(
self,
@@ -83,7 +83,7 @@ class DiscordButton(BaseMessageComponent):
class DiscordReference(BaseMessageComponent):
"""Discord引用组件"""
type: str = "discord_reference"
type: ComponentType = ComponentType.DiscordReference
def __init__(self, message_id: str, channel_id: str) -> None:
self.message_id = message_id
@@ -93,7 +93,7 @@ class DiscordReference(BaseMessageComponent):
class DiscordView(BaseMessageComponent):
"""Discord视图组件,包含按钮和选择菜单"""
type: str = "discord_view"
type: ComponentType = ComponentType.DiscordView
def __init__(
self,
@@ -117,7 +117,7 @@ class DiscordView(BaseMessageComponent):
if component.url:
# URL按钮
button = discord.ui.Button(
button: discord.ui.Button[discord.ui.View] = discord.ui.Button(
label=component.label,
style=discord.ButtonStyle.link,
url=component.url,
@@ -51,7 +51,7 @@ class DiscordPlatformAdapter(Platform):
super().__init__(platform_config, event_queue)
self.settings = platform_settings
self.client_self_id: str | None = None
self.registered_handlers = []
self.registered_handlers: list[Any] = []
self.sdk_plugin_bridge = None
# 指令注册相关
self.enable_command_register = self.config.get("discord_command_register", True)
@@ -230,7 +230,7 @@ class DiscordPlatformAdapter(Platform):
user_id=str(message.author.id),
nickname=message.author.display_name,
)
message_chain = []
message_chain: list[Any] = []
# 如果机器人被 @,在 message_chain 开头添加 At 组件
if self.client and self.client.user and bot_was_mentioned:
message_chain.insert(
@@ -46,7 +46,7 @@ class KookClient:
# 状态/计算字段
self.running = False
self.session_id = None
self.session_id: str | None = None
self.last_sn = 0 # 记录最后处理的消息序号
self.last_heartbeat_time = 0
self.heartbeat_failed_count = 0
@@ -237,9 +237,7 @@ class LarkMessageEvent(AstrMessageEvent):
def _open_image():
return open(file_path, "rb")
image_file = await asyncio.to_thread(
lambda: open(file_path, "rb")
)
image_file = await asyncio.to_thread(_open_image)
except Exception as e:
logger.error(f"[Lark] 无法打开图片文件: {e}")
continue
@@ -96,7 +96,7 @@ class MisskeyPlatformAdapter(Platform):
self._running = False
self.client_self_id = ""
self._bot_username = ""
self._user_cache = {}
self._user_cache: dict[str, Any] = {}
def meta(self) -> PlatformMetadata:
default_config = {
@@ -200,8 +200,7 @@ class MisskeyPlatformAdapter(Platform):
try:
if not isinstance(message.raw_message, dict):
message.raw_message = {}
raw_msg: dict[str, Any] = message.raw_message
raw_msg["poll"] = poll
message.raw_message["poll"] = poll
message.__setattr__("poll", poll)
except Exception:
pass
@@ -579,7 +578,7 @@ class MisskeyPlatformAdapter(Platform):
if fallback_urls:
appended = "\n" + "\n".join(fallback_urls)
text = (text or "") + appended
payload: dict[str, Any] = {"toUserId": user_id, "text": text}
payload = {"toUserId": user_id, "text": text}
if file_ids:
# 聊天消息只支持单个文件,使用 fileId 而不是 fileIds
payload["fileId"] = file_ids[0]
@@ -379,7 +379,7 @@ def process_at_mention(
client_self_id: str,
) -> tuple[list[str], str]:
"""处理@提及逻辑,返回消息部分列表和处理后的文本"""
message_parts = []
message_parts: list[str] = []
if not raw_text:
return message_parts, ""
@@ -281,7 +281,7 @@ class QQOfficialMessageEvent(AstrMessageEvent):
payload["content"] = plain_text or None
ret = await self._send_with_markdown_fallback(
send_func=lambda retry_payload: self.bot.api.post_group_message(
group_openid=source.group_openid,
group_openid=source.group_openid or "",
**retry_payload,
),
payload=payload,
@@ -577,7 +577,7 @@ class QQOfficialMessageEvent(AstrMessageEvent):
logger.error(f"[QQOfficial] post_c2c_message: 响应不是 dict: {result}")
return None
return message.Message(**cast(dict[str, Any], result))
return message.Message(**cast(dict[str, Any], result)) # type: ignore[typeddict-item]
@staticmethod
async def _parse_to_qqofficial(message: MessageChain):
@@ -10,6 +10,7 @@ from pathlib import Path
from types import SimpleNamespace
from typing import Any, cast
import anyio
import botpy
import botpy.message
from botpy import Client
@@ -344,8 +345,8 @@ class QQOfficialPlatformAdapter(Platform):
url: str,
filename: str,
) -> Record:
temp_dir = Path(get_astrbot_temp_path())
temp_dir.mkdir(parents=True, exist_ok=True)
temp_dir = anyio.Path(get_astrbot_temp_path())
await temp_dir.mkdir(parents=True, exist_ok=True)
ext = Path(filename).suffix.lower()
source_ext = ext or ".audio"
@@ -1,6 +1,7 @@
import asyncio
import json
import time
from typing import Any
from xml.etree import ElementTree as ET
import websockets
@@ -64,7 +65,7 @@ class SatoriPlatformAdapter(Platform):
self.ws: ClientConnection | None = None
self.session: ClientSession | None = None
self.sequence = 0
self.logins = []
self.logins: list[Any] = []
self.running = False
self.heartbeat_task: asyncio.Task | None = None
self.ready_received = False
@@ -188,7 +189,7 @@ class SatoriPlatformAdapter(Platform):
if self._is_websocket_closed(self.ws):
raise Exception("WebSocket连接已关闭")
identify_payload = {
identify_payload: dict[str, Any] = {
"op": 3, # IDENTIFY
"body": {
"token": str(self.token) if self.token else "", # 字符串
@@ -591,7 +592,7 @@ class SatoriPlatformAdapter(Platform):
async def parse_satori_elements(self, content: str) -> list:
"""解析 Satori 消息元素"""
elements = []
elements: list[Any] = []
if not content:
return elements
@@ -162,7 +162,7 @@ class SatoriPlatformEvent(AstrMessageEvent):
async def send_streaming(self, generator, use_fallback: bool = False):
try:
content_parts = []
content_parts: list[str] = []
async for chain in generator:
if isinstance(chain, MessageChain):
@@ -228,7 +228,7 @@ class TelegramPlatformAdapter(Platform):
def collect_commands(self) -> list[BotCommand]:
"""从注册的处理器中收集所有指令"""
command_dict = {}
command_dict: dict[str, BotCommand] = {}
skip_commands = {"start"}
for handler_md in star_handlers_registry:
@@ -10,6 +10,7 @@ import anyio
from astrbot.core.db.po import Attachment
from astrbot.core.message.components import (
BaseMessageComponent,
File,
Image,
Json,
@@ -59,7 +60,7 @@ async def parse_webchat_message_parts(
tuple[list, list[str], bool]:
(components, plain_text_parts, has_non_reply_content)
"""
components = []
components: list[BaseMessageComponent] = []
text_parts: list[str] = []
has_content = False
@@ -243,7 +244,7 @@ def webchat_message_parts_to_message_chain(
*,
strict: bool = False,
) -> MessageChain:
components = []
components: list[BaseMessageComponent] = []
has_content = False
for part in message_parts:
@@ -372,8 +373,8 @@ async def message_chain_to_storage_message_parts(
insert_attachment: AttachmentInserter,
attachments_dir: str | Path,
) -> list[dict]:
target_dir = anyio.Path(attachments_dir)
await target_dir.mkdir(parents=True, exist_ok=True)
target_dir = Path(attachments_dir)
await anyio.Path(target_dir).mkdir(parents=True, exist_ok=True)
parts: list[dict] = []
for comp in message_chain.chain:
@@ -7,7 +7,6 @@ from typing import Any, cast
import aiofiles
import quart
from requests import Response
from wechatpy.enterprise import WeChatClient, parse_message
from wechatpy.enterprise.crypto import WeChatCrypto
from wechatpy.enterprise.messages import ImageMessage, TextMessage, VoiceMessage
@@ -340,7 +339,7 @@ class WecomPlatformAdapter(Platform):
abm.session_id = abm.sender.user_id
abm.raw_message = msg
elif isinstance(msg, VoiceMessage):
resp: Response = await asyncio.get_running_loop().run_in_executor(
resp = await asyncio.get_running_loop().run_in_executor(
None,
self.client.media.download,
msg.media_id,
@@ -396,7 +395,7 @@ class WecomPlatformAdapter(Platform):
abm.message_str = text
elif msgtype == "image":
media_id = msg.get("image", {}).get("media_id", "")
resp: Response = await asyncio.get_running_loop().run_in_executor(
resp = await asyncio.get_running_loop().run_in_executor(
None,
self.client.media.download,
media_id,
@@ -408,7 +407,7 @@ class WecomPlatformAdapter(Platform):
abm.message = [Image(file=path, url=path)]
elif msgtype == "voice":
media_id = msg.get("voice", {}).get("media_id", "")
resp: Response = await asyncio.get_running_loop().run_in_executor(
resp = await asyncio.get_running_loop().run_in_executor(
None,
self.client.media.download,
media_id,
@@ -179,7 +179,7 @@ class WecomAIBotAdapter(Platform):
except Exception as e:
logger.error(f"处理队列消息时发生异常: {e}")
async def _process_message( # type: ignore[invalid-method-override]
async def _process_message(
self,
message_data: dict[str, Any],
callback_params: dict[str, str],
@@ -356,7 +356,7 @@ class WecomAIBotAdapter(Platform):
logger.error("处理欢迎消息时发生异常: %s", e)
return None
async def _process_long_connection_payload( # type: ignore[invalid-method-override]
async def _process_long_connection_payload(
self,
payload: dict[str, Any],
) -> None:
@@ -425,7 +425,7 @@ class WecomAIBotAdapter(Platform):
},
)
async def _send_long_connection_respond_msg( # type: ignore[invalid-method-override]
async def _send_long_connection_respond_msg(
self,
req_id: str,
body: dict[str, Any],
@@ -451,7 +451,7 @@ class WecomAIBotAdapter(Platform):
user_id = message_data.get("from", {}).get("userid", "default_user")
return format_session_id("wecomai", user_id)
async def _enqueue_message( # type: ignore[invalid-method-override]
async def _enqueue_message(
self,
message_data: dict[str, Any],
callback_params: dict[str, str],
@@ -482,7 +482,7 @@ class WecomAIBotAdapter(Platform):
image_base64 = []
_img_url_to_process: list[tuple[str, str | None]] = []
msg_items = []
msg_items: list[dict[str, Any]] = []
if msgtype == WecomAIBotConstants.MSG_TYPE_TEXT:
content = WecomAIBotMessageParser.parse_text_message(message_data)
@@ -561,7 +561,7 @@ class WecomAIBotAdapter(Platform):
logger.debug(f"WecomAIAdapter: {abm.message}")
return abm
async def send_by_session( # type: ignore[invalid-method-override]
async def send_by_session(
self,
session: MessageSesion,
message_chain: MessageChain,
@@ -585,9 +585,9 @@ class WecomAIBotAdapter(Platform):
)
await super().send_by_session(session, message_chain)
def run( # type: ignore[invalid-method-override]
def run(
self,
) -> Awaitable[Any]:
) -> Coroutine[Any, Any, None]:
"""运行适配器,同时启动HTTP服务器和队列监听器"""
async def run_both() -> None:
@@ -95,17 +95,18 @@ class WeixinOCClient:
def _build_media_cipher(key: bytes):
# Weixin OC CDN media transport only exchanges an `aeskey`; no IV is
# negotiated by the upstream API, so ECB is required for compatibility.
# codeql[py/weak-cryptographic-algorithm]
# This is a known limitation of the WeChat API design.
# nosec[ECB]
return AES.new(key, AES.MODE_ECB)
@classmethod
def encrypt_cdn_payload(cls, data: bytes, key: bytes) -> bytes:
cipher = cls._build_media_cipher(key)
cipher = cls._build_media_cipher(key) # nosec[ECB]
return cipher.encrypt(cls.pkcs7_pad(data))
@classmethod
def decrypt_cdn_payload(cls, encrypted: bytes, key: bytes) -> bytes:
cipher = cls._build_media_cipher(key)
cipher = cls._build_media_cipher(key) # nosec[ECB]
return cls.pkcs7_unpad(cipher.decrypt(encrypted))
@staticmethod
@@ -29,7 +29,7 @@ class WeixinOCMessageEvent(AstrMessageEvent):
platform: WeixinOCAdapter,
) -> None:
super().__init__(message_str, message_obj, platform_meta, session_id)
self.platform = platform
self.adapter = platform
self._typing_owner_id: str | None = None
def _get_typing_owner_id(self) -> str:
@@ -62,17 +62,17 @@ class WeixinOCMessageEvent(AstrMessageEvent):
async def send(self, message: MessageChain) -> None:
if not message.chain:
return
await self.platform.send_by_session(self.session, message)
await self.adapter.send_by_session(self.session, message)
await super().send(message)
async def send_typing(self) -> None:
await self.platform.start_typing(
await self.adapter.start_typing(
self.session.session_id,
self._get_typing_owner_id(),
)
async def stop_typing(self) -> None:
await self.platform.stop_typing(
await self.adapter.stop_typing(
self.session.session_id,
self._get_typing_owner_id(),
)