Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 15 additions & 1 deletion astrbot/builtin_stars/astrbot/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,8 +114,22 @@ async def empty_mention_waiter(
controller: SessionController,
event: AstrMessageEvent,
) -> None:
# Process both text and non-text messages (e.g. pure images).
# Previously, empty message_str caused an early return that
# silently dropped the event after stop_event() had already
# been called by handle_session_control_agent.
if not event.message_str or not event.message_str.strip():
return
# Degenerate case: completely empty message chain —
# do not re-queue, as it would cause an infinite loop.
# Stop the controller so the waiter session ends cleanly
# and does not linger to intercept subsequent messages.
if not event.get_messages():
logger.warning(
"empty_mention_waiter: received event with "
"empty message_str and empty message chain, skipping"
)
controller.stop()
return
event.message_obj.message.insert(
0,
Comp.At(qq=event.get_self_id(), name=event.get_self_id()),
Expand Down
65 changes: 54 additions & 11 deletions astrbot/core/utils/session_waiter.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from typing import Any

import astrbot.core.message.components as Comp
from astrbot.core import logger
from astrbot.core.platform import AstrMessageEvent

USER_SESSIONS: dict[str, "SessionWaiter"] = {} # 存储 SessionWaiter 实例
Expand Down Expand Up @@ -87,18 +88,40 @@ def get_history_chains(self) -> list[list[Comp.BaseMessageComponent]]:
return self.history_chains


class SessionFilter:
"""如何界定一个会话"""
class SessionFilter(abc.ABC):
"""How to identify a session scope."""

@abc.abstractmethod
def filter(self, event: AstrMessageEvent) -> str:
"""根据事件返回一个会话标识符"""
"""Return a session identifier derived from the event."""


class DefaultSessionFilter(SessionFilter):
def filter(self, event: AstrMessageEvent) -> str:
"""默认实现,返回统一消息来源字符串作为会话标识符"""
return event.unified_msg_origin
"""Return a composite session key from both conversation and sender.

The key format is ``"{unified_msg_origin}:{sender_id}"`` so that two
distinct senders in the same group chat produce different keys. When
the sender id cannot be determined (empty string), the key uses the
placeholder ``"<unknown>"`` as the sender component
(``f"{umo}:<unknown>"``) to avoid cross-user key collision, and a
warning is logged.

Args:
event: The incoming message event.

Returns:
A session identifier scoped to both the conversation and the sender.
"""
sender_id = event.get_sender_id()
if sender_id:
return f"{event.unified_msg_origin}:{sender_id}"
logger.warning(
"session_waiter: sender_id is empty for event from %s, "
"using '<unknown>' placeholder to avoid cross-user key collision",
event.unified_msg_origin,
)
return f"{event.unified_msg_origin}:<unknown>"


class SessionWaiter:
Expand Down Expand Up @@ -128,6 +151,12 @@ async def register_wait(
) -> Any:
"""等待外部输入并处理"""
self.handler = handler
existing = USER_SESSIONS.get(self.session_id)
if existing is not None and existing is not self:
logger.warning(
"session_waiter: overwriting existing waiter for session %s",
self.session_id,
)
USER_SESSIONS[self.session_id] = self

# 开始一个会话保持事件
Expand All @@ -142,12 +171,26 @@ async def register_wait(
self._cleanup()

def _cleanup(self, error: Exception | None = None) -> None:
"""清理会话"""
USER_SESSIONS.pop(self.session_id, None)
try:
FILTERS.remove(self.session_filter)
except ValueError:
pass
"""清理会话。

Only removes this waiter from ``USER_SESSIONS`` if the stored instance
is *self* — a newer waiter registered under the same key must not be
evicted by a stale cleanup.
"""
stored = USER_SESSIONS.get(self.session_id)
if stored is self:
USER_SESSIONS.pop(self.session_id, None)
elif stored is not None:
logger.warning(
"session_waiter: skipping _cleanup for session %s — "
"a newer waiter is already registered",
self.session_id,
)
# Use identity check to avoid removing a newer waiter's filter.
for i, f in enumerate(FILTERS):
if f is self.session_filter:
FILTERS.pop(i)
break
self.session_controller.stop(error)

@classmethod
Expand Down
Loading