diff --git a/astrbot/core/platform/sources/qqofficial/qqofficial_platform_adapter.py b/astrbot/core/platform/sources/qqofficial/qqofficial_platform_adapter.py index 8f73ccae88..6e74dec6b3 100644 --- a/astrbot/core/platform/sources/qqofficial/qqofficial_platform_adapter.py +++ b/astrbot/core/platform/sources/qqofficial/qqofficial_platform_adapter.py @@ -27,6 +27,7 @@ Platform, PlatformMetadata, ) +from astrbot.core import sp from astrbot.core.message.components import BaseMessageComponent from astrbot.core.platform.astr_message_event import MessageSesion from astrbot.core.utils.media_utils import MediaResolver @@ -204,7 +205,7 @@ async def on_group_at_message_create( ) abm.group_id = cast(str, message.group_openid) abm.session_id = abm.group_id - self.platform.remember_session_scene(abm.session_id, "group") + await self.platform.remember_session_scene(abm.session_id, "group") self._commit(abm) async def on_group_message_create( @@ -216,7 +217,7 @@ async def on_group_message_create( ) abm.group_id = cast(str, message.group_openid) abm.session_id = abm.group_id - self.platform.remember_session_scene(abm.session_id, "group") + await self.platform.remember_session_scene(abm.session_id, "group") self._commit(abm) # 收到频道消息 @@ -227,7 +228,7 @@ async def on_at_message_create(self, message: botpy.message.Message) -> None: ) abm.group_id = message.channel_id abm.session_id = abm.group_id - self.platform.remember_session_scene(abm.session_id, "channel") + await self.platform.remember_session_scene(abm.session_id, "channel") self._commit(abm) # 收到私聊消息 @@ -239,7 +240,7 @@ async def on_direct_message_create( MessageType.FRIEND_MESSAGE, ) abm.session_id = abm.sender.user_id - self.platform.remember_session_scene(abm.session_id, "friend") + await self.platform.remember_session_scene(abm.session_id, "friend") self._commit(abm) # 收到 C2C 消息 @@ -249,7 +250,7 @@ async def on_c2c_message_create(self, message: botpy.message.C2CMessage) -> None MessageType.FRIEND_MESSAGE, ) abm.session_id = abm.sender.user_id - self.platform.remember_session_scene(abm.session_id, "friend") + await self.platform.remember_session_scene(abm.session_id, "friend") self._commit(abm) def _commit(self, abm: AstrBotMessage) -> None: @@ -321,7 +322,6 @@ def __init__( self._session_last_message_id: dict[str, str] = {} self._session_scene: dict[str, str] = {} - self._allow_group_proactive_send = True self.test_mode = os.environ.get("TEST_MODE", "off") == "on" @@ -345,6 +345,9 @@ async def _send_by_session_common( Returns: None. + + Raises: + ValueError: The group or channel delivery scene is unknown. """ if session.message_type == MessageType.GROUP_MESSAGE: session = MessageSesion( @@ -380,24 +383,22 @@ async def _send_by_session_common( ): return - # 主动推送不需要 msg_id,见 https://github.com/AstrBotDevs/AstrBot/issues/7904 - msg_id = self._session_last_message_id.get(session.session_id) scene = self._session_scene.get(session.session_id) - allow_group_proactive_send = ( - session.message_type == MessageType.GROUP_MESSAGE - and scene == "group" - and getattr(self, "_allow_group_proactive_send", False) - ) - if ( - not msg_id - and session.message_type != MessageType.FRIEND_MESSAGE - and not allow_group_proactive_send - ): - logger.warning( - "[QQOfficial] No cached msg_id for session: %s, skip send_by_session", - session.session_id, - ) - return + if session.message_type == MessageType.GROUP_MESSAGE: + if scene is None: + scene = await sp.get_async( + "qqofficial", + f"{self.meta().id}:{self.appid}", + f"group_scene:{session.session_id}", + None, + ) + if scene not in ("group", "channel"): + raise ValueError( + "[QQOfficial] Unknown delivery scene for session " + f"{session.session_id}; receive a message from this group or " + "channel first to establish its route." + ) + self._session_scene[session.session_id] = scene use_md = getattr(message_chain, "use_markdown_", None) if use_md is False or (use_md is None and not self.use_markdown_default): @@ -407,8 +408,7 @@ async def _send_by_session_common( "markdown": MarkdownPayload(content=plain_text) if plain_text else None, "msg_type": 2, } - if msg_id and not allow_group_proactive_send: - payload["msg_id"] = msg_id + # Session sends are proactive; cached reply IDs may be expired. ret: Any = None send_helper = SimpleNamespace(bot=self.client) @@ -567,9 +567,31 @@ def remember_session_message_id(self, session_id: str, message_id: str) -> None: return self._session_last_message_id[session_id] = message_id - def remember_session_scene(self, session_id: str, scene: str) -> None: + async def remember_session_scene(self, session_id: str, scene: str) -> None: + """Persist group delivery routes before dispatching incoming events. + + Args: + session_id: Raw group OpenID, channel ID, or private sender ID. + scene: Delivery scene reported by the incoming event. + """ if not session_id or not scene: return + if self._session_scene.get(session_id) == scene: + return + if scene in ("group", "channel"): + try: + await sp.put_async( + "qqofficial", + f"{self.meta().id}:{self.appid}", + f"group_scene:{session_id}", + scene, + ) + except Exception: + logger.exception( + "[QQOfficial] Failed to persist delivery scene for session %s", + session_id, + ) + raise self._session_scene[session_id] = scene def _extract_message_id(self, ret: Any) -> str | None: 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 b3c1ea9cf9..d7fb3e631c 100644 --- a/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_adapter.py +++ b/astrbot/core/platform/sources/qqofficial_webhook/qo_webhook_adapter.py @@ -41,7 +41,7 @@ async def on_group_at_message_create( ) abm.group_id = cast(str, message.group_openid) abm.session_id = abm.group_id - self.platform.remember_session_scene(abm.session_id, "group") + await self.platform.remember_session_scene(abm.session_id, "group") self._commit(abm) async def on_group_message_create( @@ -53,7 +53,7 @@ async def on_group_message_create( ) abm.group_id = cast(str, message.group_openid) abm.session_id = abm.group_id - self.platform.remember_session_scene(abm.session_id, "group") + await self.platform.remember_session_scene(abm.session_id, "group") self._commit(abm) # 收到频道消息 @@ -64,7 +64,7 @@ async def on_at_message_create(self, message: botpy.message.Message) -> None: ) abm.group_id = message.channel_id abm.session_id = abm.group_id - self.platform.remember_session_scene(abm.session_id, "channel") + await self.platform.remember_session_scene(abm.session_id, "channel") self._commit(abm) # 收到私聊消息 @@ -76,7 +76,7 @@ async def on_direct_message_create( MessageType.FRIEND_MESSAGE, ) abm.session_id = abm.sender.user_id - self.platform.remember_session_scene(abm.session_id, "friend") + await self.platform.remember_session_scene(abm.session_id, "friend") self._commit(abm) # 收到 C2C 消息 @@ -86,7 +86,7 @@ async def on_c2c_message_create(self, message: botpy.message.C2CMessage) -> None MessageType.FRIEND_MESSAGE, ) abm.session_id = abm.sender.user_id - self.platform.remember_session_scene(abm.session_id, "friend") + await self.platform.remember_session_scene(abm.session_id, "friend") self._commit(abm) def _commit(self, abm: AstrBotMessage) -> None: @@ -125,7 +125,6 @@ def __init__( self.webhook_helper = None self._session_last_message_id: dict[str, str] = {} self._session_scene: dict[str, str] = {} - self._allow_group_proactive_send = True async def send_by_session( self, @@ -143,10 +142,16 @@ def remember_session_message_id(self, session_id: str, message_id: str) -> None: return self._session_last_message_id[session_id] = message_id - def remember_session_scene(self, session_id: str, scene: str) -> None: - if not session_id or not scene: - return - self._session_scene[session_id] = scene + async def remember_session_scene(self, session_id: str, scene: str) -> None: + """Persist a delivery route using the shared QQ Official implementation. + + Args: + session_id: Raw destination ID reported by QQ. + scene: Delivery scene reported by the incoming event. + """ + await QQOfficialPlatformAdapter.remember_session_scene( + cast(Any, self), session_id, scene + ) def _extract_message_id(self, ret: Any) -> str | None: if isinstance(ret, dict): diff --git a/astrbot/core/star/context.py b/astrbot/core/star/context.py index b409bbeaba..00f9586ba1 100644 --- a/astrbot/core/star/context.py +++ b/astrbot/core/star/context.py @@ -626,11 +626,11 @@ async def send_message( 是否找到匹配的平台。 Raises: - ValueError: session 字符串不合法时抛出。 + ValueError: The session string is invalid or the adapter cannot resolve + its delivery route. Note: 当 session 为字符串时,会尝试解析为 MessageSession 对象。(类名为MessageSesion是因为历史遗留拼写错误) - qq_official(QQ 官方 API 平台) 不支持此方法。 """ if isinstance(session, str): try: diff --git a/tests/test_qqofficial_group_message_create.py b/tests/test_qqofficial_group_message_create.py index 9e7f3047dd..3d53c9e4e2 100644 --- a/tests/test_qqofficial_group_message_create.py +++ b/tests/test_qqofficial_group_message_create.py @@ -226,7 +226,7 @@ async def test_group_message_create_handler_maps_group_session_and_scene(): remembered_ids: list[tuple[str, str]] = [] class PlatformStub: - def remember_session_scene(self, session_id: str, scene: str) -> None: + async def remember_session_scene(self, session_id: str, scene: str) -> None: remembered_scenes.append((session_id, scene)) def remember_session_message_id(self, session_id: str, message_id: str) -> None: diff --git a/tests/test_qqofficial_proactive_routing.py b/tests/test_qqofficial_proactive_routing.py new file mode 100644 index 0000000000..ce43aa55cc --- /dev/null +++ b/tests/test_qqofficial_proactive_routing.py @@ -0,0 +1,214 @@ +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock + +import pytest + +from astrbot.api.event import MessageChain +from astrbot.api.message_components import Plain +from astrbot.core.db.sqlite import SQLiteDatabase +from astrbot.core.platform.message_session import MessageSession +from astrbot.core.platform.message_type import MessageType +from astrbot.core.platform.sources.qqofficial import qqofficial_platform_adapter as qq +from astrbot.core.platform.sources.qqofficial_webhook.qo_webhook_adapter import ( + QQOfficialWebhookPlatformAdapter, +) +from astrbot.core.utils.shared_preferences import SharedPreferences + + +@pytest.fixture(params=[qq.QQOfficialPlatformAdapter, QQOfficialWebhookPlatformAdapter]) +def adapter_factory(request): + def create(**overrides): + config = { + "id": "qq-test", + "appid": "123", + "secret": "secret", + "enable_group_c2c": True, + "enable_guild_direct_message": False, + "use_markdown": False, + **overrides, + } + adapter = request.param(config, {}, asyncio.Queue()) + adapter.client.api = SimpleNamespace( + post_group_message=AsyncMock(return_value={"id": "sent-group"}), + post_message=AsyncMock(return_value={"id": "sent-channel"}), + ) + return adapter + + return create + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "handler,scene", + [ + ("on_group_at_message_create", "group"), + ("on_group_message_create", "group"), + ("on_at_message_create", "channel"), + ], +) +async def test_incoming_event_waits_for_route_persistence( + adapter_factory, handler, scene, monkeypatch +): + entered = asyncio.Event() + release = asyncio.Event() + + async def write(*args): + entered.set() + await release.wait() + + store = SimpleNamespace(put_async=AsyncMock(side_effect=write)) + monkeypatch.setattr(qq, "sp", store) + monkeypatch.setattr( + qq.QQOfficialPlatformAdapter, + "_parse_from_qqofficial", + AsyncMock(return_value=SimpleNamespace()), + ) + adapter = adapter_factory() + commit = Mock() + monkeypatch.setattr(adapter.client, "_commit", commit) + task = asyncio.create_task( + getattr(adapter.client, handler)( + SimpleNamespace(group_openid="target", channel_id="target") + ) + ) + try: + await asyncio.wait_for(entered.wait(), timeout=5) + commit.assert_not_called() + release.set() + await asyncio.wait_for(task, timeout=5) + commit.assert_called_once() + store.put_async.assert_awaited_once_with( + "qqofficial", "qq-test:123", "group_scene:target", scene + ) + finally: + release.set() + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("scene", ["group", "channel"]) +async def test_route_survives_new_adapter_and_store( + adapter_factory, scene, tmp_path, monkeypatch +): + database = SQLiteDatabase(str(tmp_path / "routes.db")) + await database.initialize() + store = SharedPreferences(database, tmp_path / "unused.json") + try: + monkeypatch.setattr(qq, "sp", store) + original = adapter_factory() + await original.remember_session_scene("target", scene) + await store.close() + await database.engine.dispose() + + database = SQLiteDatabase(str(tmp_path / "routes.db")) + await database.initialize() + store = SharedPreferences(database, tmp_path / "unused.json") + monkeypatch.setattr(qq, "sp", store) + restarted = adapter_factory() + assert restarted._session_scene == {} + assert restarted._session_last_message_id == {} + + # Unique sessions must still resolve to the raw destination ID. + await restarted.send_by_session( + MessageSession("qq-test", MessageType.GROUP_MESSAGE, "member_target"), + MessageChain(chain=[Plain("scheduled notification")]), + ) + api = restarted.client.api + selected = api.post_group_message if scene == "group" else api.post_message + other = api.post_message if scene == "group" else api.post_group_message + selected.assert_awaited_once() + other.assert_not_awaited() + payload = selected.await_args.kwargs + assert payload["group_openid" if scene == "group" else "channel_id"] == "target" + assert "msg_id" not in payload + assert restarted._session_scene == {"target": scene} + + # Neither another instance nor another bot may reuse this route. + for overrides in ({"id": "other-instance"}, {"appid": "other-app"}): + unrelated = adapter_factory(**overrides) + with pytest.raises(ValueError, match="Unknown delivery scene"): + await unrelated.send_by_session( + MessageSession( + unrelated.meta().id, MessageType.GROUP_MESSAGE, "target" + ), + MessageChain(chain=[Plain("hello")]), + ) + unrelated.client.api.post_group_message.assert_not_awaited() + unrelated.client.api.post_message.assert_not_awaited() + finally: + await store.close() + await database.engine.dispose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("scene", ["group", "channel"]) +@pytest.mark.parametrize("cached_id", [None, "expired-reply-id"]) +async def test_proactive_send_never_uses_cached_reply_id( + adapter_factory, scene, cached_id, monkeypatch +): + store = SimpleNamespace(get_async=AsyncMock()) + monkeypatch.setattr(qq, "sp", store) + adapter = adapter_factory() + adapter._session_scene["target"] = scene + if cached_id: + adapter._session_last_message_id["target"] = cached_id + session = MessageSession("qq-test", MessageType.GROUP_MESSAGE, "target") + # A second send must not use the bot's own ID cached by the first send. + for _ in range(2): + await adapter.send_by_session(session, MessageChain(chain=[Plain("hello")])) + api = adapter.client.api + selected = api.post_group_message if scene == "group" else api.post_message + assert selected.await_count == 2 + for call in selected.await_args_list: + assert "msg_id" not in call.kwargs + if scene == "channel": + assert "msg_type" not in call.kwargs + store.get_async.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "stored_scene", [None, "friend", "invalid", {"scene": "group"}] +) +async def test_unknown_route_never_calls_a_send_api( + adapter_factory, stored_scene, monkeypatch +): + monkeypatch.setattr( + qq, "sp", SimpleNamespace(get_async=AsyncMock(return_value=stored_scene)) + ) + adapter = adapter_factory() + adapter._session_last_message_id["target"] = "cached-id" + with pytest.raises(ValueError, match="Unknown delivery scene"): + await adapter.send_by_session( + MessageSession("qq-test", MessageType.GROUP_MESSAGE, "target"), + MessageChain(chain=[Plain("hello")]), + ) + adapter.client.api.post_group_message.assert_not_awaited() + adapter.client.api.post_message.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_route_write_failure_is_retried_and_unchanged_routes_are_not_rewritten( + adapter_factory, monkeypatch +): + store = SimpleNamespace( + put_async=AsyncMock(side_effect=[OSError("disk error"), None, None]) + ) + monkeypatch.setattr(qq, "sp", store) + adapter = adapter_factory() + await adapter.remember_session_scene("target", "group") + assert "target" not in adapter._session_scene + await adapter.remember_session_scene("target", "group") + assert adapter._session_scene["target"] == "group" + await adapter.remember_session_scene("target", "group") + assert store.put_async.await_count == 2 + await adapter.remember_session_scene("target", "channel") + assert adapter._session_scene["target"] == "channel" + store.put_async.assert_awaited_with( + "qqofficial", "qq-test:123", "group_scene:target", "channel" + ) + await adapter.remember_session_scene("private-user", "friend") + assert store.put_async.await_count == 3