mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
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:
@@ -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(),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user