From 0c674dd2b0618d847ef5497c4429338ce811a8c3 Mon Sep 17 00:00:00 2001 From: icywwwaa <3223850276@qq.com> Date: Thu, 20 Aug 2026 23:23:48 +0800 Subject: [PATCH] fix: scope the default session waiter identity to the sender DefaultSessionFilter keyed sessions on unified_msg_origin alone, which is identical for every member of a group chat. Any member's next message therefore hit a waiter another member had registered: the message was stopped, given a synthetic At component and re-dispatched, so the bot replied to the wrong user with the wrong context. Keying on unified_msg_origin plus sender_id closes that hole without reintroducing #1326, where a sender_id-only key let one user's waiter capture their messages in other sessions. This also restores the behaviour the session-control guide documents. Fixes #9377 --- astrbot/core/utils/session_waiter.py | 16 +++- tests/unit/test_session_waiter.py | 135 +++++++++++++++++++++++++++ 2 files changed, 149 insertions(+), 2 deletions(-) create mode 100644 tests/unit/test_session_waiter.py diff --git a/astrbot/core/utils/session_waiter.py b/astrbot/core/utils/session_waiter.py index b327a61843..1dde84c69e 100644 --- a/astrbot/core/utils/session_waiter.py +++ b/astrbot/core/utils/session_waiter.py @@ -97,8 +97,20 @@ def filter(self, event: AstrMessageEvent) -> str: class DefaultSessionFilter(SessionFilter): def filter(self, event: AstrMessageEvent) -> str: - """默认实现,返回统一消息来源字符串作为会话标识符""" - return event.unified_msg_origin + """默认实现,返回「消息来源 + 发送人」作为会话标识符。 + + 两部分都是必需的: 只用 ``unified_msg_origin`` 会让群内任意成员的下一条 + 消息命中别人注册的等待器(等待器会截获并重新投递该消息); 只用 + ``sender_id`` 又会让同一用户在其他群聊/私聊中的消息命中此等待器。 + 需要整群共享一个会话时,请自定义 :class:`SessionFilter`。 + + Args: + event: 待判定的消息事件。 + + Returns: + 会话标识符,格式为 ``{unified_msg_origin}!{sender_id}``。 + """ + return f"{event.unified_msg_origin}!{event.get_sender_id()}" class SessionWaiter: diff --git a/tests/unit/test_session_waiter.py b/tests/unit/test_session_waiter.py new file mode 100644 index 0000000000..f82c4cd806 --- /dev/null +++ b/tests/unit/test_session_waiter.py @@ -0,0 +1,135 @@ +"""Tests for the session waiter's default session identity. + +Regression coverage for the group-chat interception bug: a waiter registered by +one group member must not be triggered by a different member of the same group, +while the same member must still be isolated across different sessions. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from astrbot.core.message.components import Plain +from astrbot.core.platform.astr_message_event import AstrMessageEvent +from astrbot.core.platform.astrbot_message import AstrBotMessage, MessageMember +from astrbot.core.platform.message_type import MessageType +from astrbot.core.platform.platform_metadata import PlatformMetadata +from astrbot.core.utils.session_waiter import ( + USER_SESSIONS, + DefaultSessionFilter, + SessionController, + SessionWaiter, + session_waiter, +) + +PLATFORM_META = PlatformMetadata( + name="aiocqhttp", + description="test platform", + id="aiocqhttp", +) + + +def make_event( + sender_id: str, + session_id: str, + message_type: MessageType = MessageType.GROUP_MESSAGE, + text: str = "hello", +) -> AstrMessageEvent: + """Build a minimal group/private message event. + + Args: + sender_id: ID of the member that sent the message. + session_id: Platform session ID (group ID for group messages). + message_type: Message type of the event. + text: Plain text payload of the message. + + Returns: + A usable ``AstrMessageEvent`` for session-identity assertions. + """ + message_obj = AstrBotMessage() + message_obj.type = message_type + message_obj.self_id = "bot" + message_obj.session_id = session_id + message_obj.message_id = "1" + message_obj.sender = MessageMember(user_id=sender_id, nickname=sender_id) + message_obj.message = [Plain(text=text)] + message_obj.message_str = text + message_obj.raw_message = None + if message_type == MessageType.GROUP_MESSAGE: + message_obj.group_id = session_id + return AstrMessageEvent( + message_str=text, + message_obj=message_obj, + platform_meta=PLATFORM_META, + session_id=session_id, + ) + + +def test_default_filter_separates_members_of_the_same_group(): + """Two members of one group must map to different session identities.""" + session_filter = DefaultSessionFilter() + event_a = make_event("member_a", "group_1") + event_b = make_event("member_b", "group_1") + + assert session_filter.filter(event_a) != session_filter.filter(event_b) + + +def test_default_filter_is_stable_for_the_same_member(): + """The same member in the same group must map to one session identity.""" + session_filter = DefaultSessionFilter() + first = make_event("member_a", "group_1", text="one") + second = make_event("member_a", "group_1", text="two") + + assert session_filter.filter(first) == session_filter.filter(second) + + +def test_default_filter_separates_sessions_of_the_same_member(): + """One member must not share a waiter across groups or private chats.""" + session_filter = DefaultSessionFilter() + in_group_1 = make_event("member_a", "group_1") + in_group_2 = make_event("member_a", "group_2") + in_private = make_event( + "member_a", + "member_a", + message_type=MessageType.FRIEND_MESSAGE, + ) + + identities = { + session_filter.filter(in_group_1), + session_filter.filter(in_group_2), + session_filter.filter(in_private), + } + assert len(identities) == 3 + + +@pytest.mark.asyncio +async def test_waiter_ignores_other_members_and_accepts_the_owner(): + """A registered waiter only fires for the member that created it.""" + USER_SESSIONS.clear() + session_filter = DefaultSessionFilter() + owner_event = make_event("member_a", "group_1", text="@bot") + other_event = make_event("member_b", "group_1", text="unrelated chatter") + + triggered: list[str] = [] + + @session_waiter(timeout=5) + async def waiter(controller: SessionController, event: AstrMessageEvent) -> None: + triggered.append(event.get_sender_id()) + controller.stop() + + waiting = asyncio.create_task(waiter(owner_event, session_filter)) + await asyncio.sleep(0) + + # A different member of the same group must not reach the waiter. + await SessionWaiter.trigger(session_filter.filter(other_event), other_event) + assert triggered == [] + + # The owner's own follow-up message must reach the waiter. + follow_up = make_event("member_a", "group_1", text="the real question") + await SessionWaiter.trigger(session_filter.filter(follow_up), follow_up) + await asyncio.wait_for(waiting, timeout=5) + + assert triggered == ["member_a"] + assert USER_SESSIONS == {}