From c610719a4410ed3343b045147f226176cc5f31ef Mon Sep 17 00:00:00 2001 From: Soulter <905617992@qq.com> Date: Sat, 22 Mar 2025 19:02:49 +0800 Subject: [PATCH] =?UTF-8?q?=E2=9C=A8=20feat:=20=E4=B8=BA=E5=90=84=E5=B9=B3?= =?UTF-8?q?=E5=8F=B0=E9=80=82=E9=85=8D=E5=99=A8=E6=94=AF=E6=8C=81=E4=BC=98?= =?UTF-8?q?=E9=9B=85=E5=85=B3=E9=97=AD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- astrbot/core/core_lifecycle.py | 12 ++++++--- .../initial_loader.py} | 9 ++++--- astrbot/core/platform/manager.py | 7 ++++- astrbot/core/platform/platform.py | 2 +- .../aiocqhttp/aiocqhttp_platform_adapter.py | 17 +++++------- .../sources/dingtalk/dingtalk_adapter.py | 27 +++++++++++++++++-- .../core/platform/sources/gewechat/client.py | 14 +--------- .../gewechat/gewechat_platform_adapter.py | 3 +-- .../platform/sources/lark/lark_adapter.py | 4 +++ .../qqofficial/qqofficial_platform_adapter.py | 5 ++++ .../qqofficial_webhook/qo_webhook_adapter.py | 6 +++++ .../qqofficial_webhook/qo_webhook_server.py | 11 +------- .../platform/sources/telegram/tg_adapter.py | 4 +++ .../sources/webchat/webchat_adapter.py | 4 +++ .../platform/sources/wecom/wecom_adapter.py | 9 +++---- astrbot/dashboard/__init__.py | 3 --- astrbot/dashboard/routes/config.py | 2 +- astrbot/dashboard/routes/stat.py | 2 +- astrbot/dashboard/routes/update.py | 3 +-- astrbot/dashboard/server.py | 22 ++++++++------- main.py | 4 +-- 21 files changed, 99 insertions(+), 71 deletions(-) rename astrbot/{dashboard/dashboard_lifecycle.py => core/initial_loader.py} (82%) delete mode 100644 astrbot/dashboard/__init__.py diff --git a/astrbot/core/core_lifecycle.py b/astrbot/core/core_lifecycle.py index e5485edb4..e52d94674 100644 --- a/astrbot/core/core_lifecycle.py +++ b/astrbot/core/core_lifecycle.py @@ -40,7 +40,6 @@ class AstrBotCoreLifecycle: else: logger.setLevel(self.astrbot_config["log_level"]) self.event_queue = Queue() - self.event_queue.closed = False self.provider_manager = ProviderManager(self.astrbot_config, self.db) @@ -81,6 +80,8 @@ class AstrBotCoreLifecycle: await self.platform_manager.initialize() """根据配置实例化各个平台适配器""" + self.dashboard_shutdown_event = asyncio.Event() + def _load(self): event_bus_task = asyncio.create_task( self.event_bus.dispatch(), name="event_bus" @@ -129,11 +130,12 @@ class AstrBotCoreLifecycle: await asyncio.gather(*self.curr_tasks, return_exceptions=True) async def stop(self): - self.event_queue.closed = True for task in self.curr_tasks: task.cancel() await self.provider_manager.terminate() + await self.platform_manager.terminate() + self.dashboard_shutdown_event.set() for task in self.curr_tasks: try: @@ -143,8 +145,10 @@ class AstrBotCoreLifecycle: except Exception as e: logger.error(f"任务 {task.get_name()} 发生错误: {e}") - def restart(self): - self.event_queue.closed = True + async def restart(self): + await self.provider_manager.terminate() + await self.platform_manager.terminate() + self.dashboard_shutdown_event.set() threading.Thread( target=self.astrbot_updator._reboot, name="restart", daemon=True ).start() diff --git a/astrbot/dashboard/dashboard_lifecycle.py b/astrbot/core/initial_loader.py similarity index 82% rename from astrbot/dashboard/dashboard_lifecycle.py rename to astrbot/core/initial_loader.py index 9c5c9138d..f91a71da3 100644 --- a/astrbot/dashboard/dashboard_lifecycle.py +++ b/astrbot/core/initial_loader.py @@ -2,17 +2,16 @@ import asyncio import traceback from astrbot.core import logger from astrbot.core.core_lifecycle import AstrBotCoreLifecycle -from .server import AstrBotDashboard from astrbot.core.db import BaseDatabase from astrbot.core import LogBroker +from astrbot.dashboard.server import AstrBotDashboard -class AstrBotDashBoardLifecycle: +class InitialLoader: def __init__(self, db: BaseDatabase, log_broker: LogBroker): self.db = db self.logger = logger self.log_broker = log_broker - self.dashboard_server = None async def start(self): core_lifecycle = AstrBotCoreLifecycle(self.log_broker, self.db) @@ -25,7 +24,9 @@ class AstrBotDashBoardLifecycle: logger.critical(traceback.format_exc()) logger.critical(f"😭 初始化 AstrBot 失败:{e} !!!") - self.dashboard_server = AstrBotDashboard(core_lifecycle, self.db) + self.dashboard_server = AstrBotDashboard( + core_lifecycle, self.db, core_lifecycle.dashboard_shutdown_event + ) task = asyncio.gather(core_task, self.dashboard_server.run()) try: diff --git a/astrbot/core/platform/manager.py b/astrbot/core/platform/manager.py index 9ca3b82d1..31a6e5e85 100644 --- a/astrbot/core/platform/manager.py +++ b/astrbot/core/platform/manager.py @@ -92,7 +92,7 @@ class PlatformManager: asyncio.create_task( self._task_wrapper( asyncio.create_task( - inst.run(), name=platform_config["id"] + "_platform" + inst.run(), name=f"platform_{platform_config['type']}_{platform_config['id']}" ) ) ) @@ -142,5 +142,10 @@ class PlatformManager: # 再启动新的实例 await self.load_platform(platform_config) + async def terminate(self): + for inst in self.platform_insts: + if getattr(inst, "terminate", None): + await inst.terminate() + def get_insts(self): return self.platform_insts diff --git a/astrbot/core/platform/platform.py b/astrbot/core/platform/platform.py index 8ed0be039..bcda23e06 100644 --- a/astrbot/core/platform/platform.py +++ b/astrbot/core/platform/platform.py @@ -25,7 +25,7 @@ class Platform(abc.ABC): """ 终止一个平台的运行实例。 """ - pass + ... @abc.abstractmethod def meta(self) -> PlatformMetadata: diff --git a/astrbot/core/platform/sources/aiocqhttp/aiocqhttp_platform_adapter.py b/astrbot/core/platform/sources/aiocqhttp/aiocqhttp_platform_adapter.py index 0d11e3c0b..e41071a56 100644 --- a/astrbot/core/platform/sources/aiocqhttp/aiocqhttp_platform_adapter.py +++ b/astrbot/core/platform/sources/aiocqhttp/aiocqhttp_platform_adapter.py @@ -43,8 +43,6 @@ class AiocqhttpAdapter(Platform): "适用于 OneBot 标准的消息平台适配器,支持反向 WebSockets。", ) - self.stop = False - self.bot = CQHttp( use_ws_reverse=True, import_name="aiocqhttp", api_timeout_sec=180 ) @@ -303,22 +301,19 @@ class AiocqhttpAdapter(Platform): for handler in logging.root.handlers[:]: logging.root.removeHandler(handler) logging.getLogger("aiocqhttp").setLevel(logging.ERROR) - + self.shutdown_event = asyncio.Event() return coro async def terminate(self): - self.stop = True - await asyncio.sleep(1) + self.shutdown_event.set() + + async def shutdown_trigger_placeholder(self): + await self.shutdown_event.wait() + logger.info("aiocqhttp 适配器已被优雅地关闭") def meta(self) -> PlatformMetadata: return self.metadata - async def shutdown_trigger_placeholder(self): - # TODO: use asyncio.Event - while not self._event_queue.closed and not self.stop: # noqa: ASYNC110 - await asyncio.sleep(1) - logger.info("aiocqhttp 适配器已关闭。") - async def handle_msg(self, message: AstrBotMessage): message_event = AiocqhttpMessageEvent( message_str=message.message_str, diff --git a/astrbot/core/platform/sources/dingtalk/dingtalk_adapter.py b/astrbot/core/platform/sources/dingtalk/dingtalk_adapter.py index b38c7d5c8..29ef88086 100644 --- a/astrbot/core/platform/sources/dingtalk/dingtalk_adapter.py +++ b/astrbot/core/platform/sources/dingtalk/dingtalk_adapter.py @@ -2,6 +2,8 @@ import asyncio import uuid import aiohttp import dingtalk_stream +import time +import threading from astrbot.api.platform import ( Platform, @@ -197,9 +199,30 @@ class DingtalkPlatformAdapter(Platform): async def run(self): # await self.client_.start() - loop = asyncio.get_event_loop() # 钉钉的 SDK 并没有实现真正的异步,start() 里面有堵塞方法。 - await loop.run_in_executor(None, lambda: asyncio.run(self.client_.start())) + def start_client(loop: asyncio.AbstractEventLoop): + try: + self._shutdown_event = threading.Event() + task = loop.create_task(self.client_.start()) + self._shutdown_event.wait() + if task.done(): + task.result() + except Exception as e: + if "Graceful shutdown" in str(e): + logger.info("钉钉适配器已被优雅地关闭") + return + logger.error(f"钉钉机器人启动失败: {e}") + + loop = asyncio.get_event_loop() + await loop.run_in_executor(None, start_client, loop) + + async def terminate(self): + def monkey_patch_close(): + raise Exception("Graceful shutdown") + + self.client_.open_connection = monkey_patch_close + await self.client_.websocket.close(code=1000, reason="Graceful shutdown") + self._shutdown_event.set() def get_client(self): return self.client diff --git a/astrbot/core/platform/sources/gewechat/client.py b/astrbot/core/platform/sources/gewechat/client.py index d2f28f09d..53ee1878e 100644 --- a/astrbot/core/platform/sources/gewechat/client.py +++ b/astrbot/core/platform/sources/gewechat/client.py @@ -72,8 +72,6 @@ class SimpleGewechatClient: self.userrealnames = {} - self.stop = False - async def get_token_id(self): """获取 Gewechat Token。""" async with aiohttp.ClientSession() as session: @@ -306,17 +304,7 @@ class SimpleGewechatClient: async def start_polling(self): threading.Thread(target=asyncio.run, args=(self._set_callback_url(),)).start() - await self.server.run_task( - host="0.0.0.0", - port=self.port, - shutdown_trigger=self._shutdown_trigger_placeholder, - ) - - async def _shutdown_trigger_placeholder(self): - # TODO: use asyncio.Event - while not self.event_queue.closed and not self.stop: # noqa: ASYNC110 - await asyncio.sleep(1) - logger.info("gewechat 适配器已关闭。") + await self.server.run_task(host="0.0.0.0", port=self.port) async def check_online(self, appid: str): """检查 APPID 对应的设备是否在线。""" diff --git a/astrbot/core/platform/sources/gewechat/gewechat_platform_adapter.py b/astrbot/core/platform/sources/gewechat/gewechat_platform_adapter.py index 3dbdbba27..9c8c4f5ed 100644 --- a/astrbot/core/platform/sources/gewechat/gewechat_platform_adapter.py +++ b/astrbot/core/platform/sources/gewechat/gewechat_platform_adapter.py @@ -64,8 +64,7 @@ class GewechatPlatformAdapter(Platform): ) async def terminate(self): - self.client.stop = True - await asyncio.sleep(1) + await self.client.server.shutdown() async def logout(self): await self.client.logout() diff --git a/astrbot/core/platform/sources/lark/lark_adapter.py b/astrbot/core/platform/sources/lark/lark_adapter.py index 1ee30c482..cbc3a45bb 100644 --- a/astrbot/core/platform/sources/lark/lark_adapter.py +++ b/astrbot/core/platform/sources/lark/lark_adapter.py @@ -185,5 +185,9 @@ class LarkPlatformAdapter(Platform): # self.client.start() await self.client._connect() + async def terminate(self): + await self.client._disconnect() + logger.info("飞书(Lark) 适配器已被优雅地关闭") + def get_client(self) -> lark.Client: return self.client diff --git a/astrbot/core/platform/sources/qqofficial/qqofficial_platform_adapter.py b/astrbot/core/platform/sources/qqofficial/qqofficial_platform_adapter.py index ae1dc2563..57bc8683f 100644 --- a/astrbot/core/platform/sources/qqofficial/qqofficial_platform_adapter.py +++ b/astrbot/core/platform/sources/qqofficial/qqofficial_platform_adapter.py @@ -17,6 +17,7 @@ from astrbot.api.platform import ( MessageType, PlatformMetadata, ) +from astrbot import logger from astrbot.api.event import MessageChain from typing import Union, List from astrbot.api.message_components import Image, Plain, At @@ -204,3 +205,7 @@ class QQOfficialPlatformAdapter(Platform): def get_client(self) -> botClient: return self.client + + async def terminate(self): + await self.client.close() + logger.info("QQ 官方机器人接口 适配器已被优雅地关闭") diff --git a/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_adapter.py b/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_adapter.py index 542233591..6ad59c67e 100644 --- a/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_adapter.py +++ b/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_adapter.py @@ -13,6 +13,7 @@ from .qo_webhook_event import QQOfficialWebhookMessageEvent from ...register import register_platform_adapter from .qo_webhook_server import QQOfficialWebhook from ..qqofficial.qqofficial_platform_adapter import QQOfficialPlatformAdapter +from astrbot import logger # remove logger handler for handler in logging.root.handlers[:]: @@ -111,3 +112,8 @@ class QQOfficialWebhookPlatformAdapter(Platform): def get_client(self) -> botClient: return self.client + + async def terminate(self): + await self.client.close() + await self.webhook_helper.server.shutdown() + logger.info("QQ 机器人官方 API 适配器已经被优雅地关闭") diff --git a/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_server.py b/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_server.py index a219e2492..681999cf0 100644 --- a/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_server.py +++ b/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_server.py @@ -99,13 +99,4 @@ class QQOfficialWebhook: logger.info( f"将在 {self.callback_server_host}:{self.port} 端口启动 QQ 官方机器人 webhook 适配器。" ) - await self.server.run_task( - host=self.callback_server_host, - port=self.port, - shutdown_trigger=self.shutdown_trigger_placeholder, - ) - - async def shutdown_trigger_placeholder(self): - while not self.event_queue.closed: # noqa: ASYNC110 - await asyncio.sleep(1) - logger.info("qq_official_webhook 适配器已关闭。") + await self.server.run_task(host=self.callback_server_host, port=self.port) diff --git a/astrbot/core/platform/sources/telegram/tg_adapter.py b/astrbot/core/platform/sources/telegram/tg_adapter.py index 944a65902..d53efd13d 100644 --- a/astrbot/core/platform/sources/telegram/tg_adapter.py +++ b/astrbot/core/platform/sources/telegram/tg_adapter.py @@ -225,3 +225,7 @@ class TelegramPlatformAdapter(Platform): def get_client(self) -> ExtBot: return self.client + + async def terminate(self): + await self.application.stop() + logger.info("Telegram 适配器已被优雅地关闭") diff --git a/astrbot/core/platform/sources/webchat/webchat_adapter.py b/astrbot/core/platform/sources/webchat/webchat_adapter.py index 12b193f53..6fa3d5c59 100644 --- a/astrbot/core/platform/sources/webchat/webchat_adapter.py +++ b/astrbot/core/platform/sources/webchat/webchat_adapter.py @@ -119,3 +119,7 @@ class WebChatAdapter(Platform): ) self.commit_event(message_event) + + async def terminate(self): + # Do nothing + pass diff --git a/astrbot/core/platform/sources/wecom/wecom_adapter.py b/astrbot/core/platform/sources/wecom/wecom_adapter.py index cef83b030..e39607919 100644 --- a/astrbot/core/platform/sources/wecom/wecom_adapter.py +++ b/astrbot/core/platform/sources/wecom/wecom_adapter.py @@ -93,13 +93,8 @@ class WecomServer: await self.server.run_task( host=self.callback_server_host, port=self.port, - shutdown_trigger=self.shutdown_trigger_placeholder, ) - async def shutdown_trigger_placeholder(self): - while not self.event_queue.closed: # noqa: ASYNC110 - await asyncio.sleep(1) - logger.info("企业微信 适配器已关闭。") @register_platform_adapter("wecom", "wecom 适配器") @@ -235,3 +230,7 @@ class WecomPlatformAdapter(Platform): def get_client(self) -> WeChatClient: return self.client + + async def terminate(self): + await self.server.server.shutdown() + logger.info("企业微信 适配器已被优雅地关闭") diff --git a/astrbot/dashboard/__init__.py b/astrbot/dashboard/__init__.py deleted file mode 100644 index cf829d4d6..000000000 --- a/astrbot/dashboard/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .dashboard_lifecycle import AstrBotDashBoardLifecycle - -__all__ = ["AstrBotDashBoardLifecycle"] diff --git a/astrbot/dashboard/routes/config.py b/astrbot/dashboard/routes/config.py index 088c999f9..04f3aa27e 100644 --- a/astrbot/dashboard/routes/config.py +++ b/astrbot/dashboard/routes/config.py @@ -321,7 +321,7 @@ class ConfigRoute(Route): async def _save_astrbot_configs(self, post_configs: dict): try: save_config(post_configs, self.config, is_core=True) - self.core_lifecycle.restart() + await self.core_lifecycle.restart() except Exception as e: raise e diff --git a/astrbot/dashboard/routes/stat.py b/astrbot/dashboard/routes/stat.py index 3a21aa2b9..b20e9dca4 100644 --- a/astrbot/dashboard/routes/stat.py +++ b/astrbot/dashboard/routes/stat.py @@ -28,7 +28,7 @@ class StatRoute(Route): self.core_lifecycle = core_lifecycle async def restart_core(self): - self.core_lifecycle.restart() + await self.core_lifecycle.restart() return Response().ok().__dict__ def format_sec(self, sec: int): diff --git a/astrbot/dashboard/routes/update.py b/astrbot/dashboard/routes/update.py index ef2d10634..e9ada18f5 100644 --- a/astrbot/dashboard/routes/update.py +++ b/astrbot/dashboard/routes/update.py @@ -95,8 +95,7 @@ class UpdateRoute(Route): logger.error(f"更新依赖失败: {e}") if reboot: - # threading.Thread(target=self.astrbot_updator._reboot, args=(2, )).start() - self.core_lifecycle.restart() + await self.core_lifecycle.restart() return ( Response() .ok(None, "更新成功,AstrBot 将在 2 秒内全量重启以应用新的代码。") diff --git a/astrbot/dashboard/server.py b/astrbot/dashboard/server.py index 6fc0651fa..45aac3cd6 100644 --- a/astrbot/dashboard/server.py +++ b/astrbot/dashboard/server.py @@ -20,7 +20,12 @@ DATAPATH = os.path.abspath( class AstrBotDashboard: - def __init__(self, core_lifecycle: AstrBotCoreLifecycle, db: BaseDatabase) -> None: + def __init__( + self, + core_lifecycle: AstrBotCoreLifecycle, + db: BaseDatabase, + shutdown_event: asyncio.Event, + ) -> None: self.core_lifecycle = core_lifecycle self.config = core_lifecycle.astrbot_config self.data_path = os.path.abspath(os.path.join(DATAPATH, "dist")) @@ -46,6 +51,8 @@ class AstrBotDashboard: self.ar = AuthRoute(self.context) self.chat_route = ChatRoute(self.context, db, core_lifecycle) + self.shutdown_event = shutdown_event + async def auth_middleware(self): if not request.path.startswith("/api"): return @@ -73,11 +80,6 @@ class AstrBotDashboard: r.status_code = 401 return r - async def shutdown_trigger_placeholder(self): - while not self.core_lifecycle.event_queue.closed: # noqa: ASYNC110 - await asyncio.sleep(1) - logger.info("管理面板已关闭。") - def check_port_in_use(self, port: int) -> bool: """ 跨平台检测端口是否被占用 @@ -166,7 +168,9 @@ class AstrBotDashboard: logger.info(display) return self.app.run_task( - host=host, - port=port, - shutdown_trigger=self.shutdown_trigger_placeholder, + host=host, port=port, shutdown_trigger=self.shutdown_trigger ) + + async def shutdown_trigger(self): + await self.shutdown_event.wait() + logger.info("AstrBot WebUI 已经被优雅地关闭") diff --git a/main.py b/main.py index e66d33fef..9937bb10f 100644 --- a/main.py +++ b/main.py @@ -2,7 +2,7 @@ import os import asyncio import sys import mimetypes -from astrbot.dashboard import AstrBotDashBoardLifecycle +from astrbot.core.initial_loader import InitialLoader from astrbot.core import db_helper from astrbot.core import logger, LogManager, LogBroker from astrbot.core.config.default import VERSION @@ -79,5 +79,5 @@ if __name__ == "__main__": # print logo logger.info(logo_tmpl) - dashboard_lifecycle = AstrBotDashBoardLifecycle(db, log_broker) + dashboard_lifecycle = InitialLoader(db, log_broker) asyncio.run(dashboard_lifecycle.start())