diff --git a/astrbot/core/core_lifecycle.py b/astrbot/core/core_lifecycle.py index db8a6ddc1b..9961a8c52b 100644 --- a/astrbot/core/core_lifecycle.py +++ b/astrbot/core/core_lifecycle.py @@ -462,7 +462,7 @@ async def load_pipeline_scheduler(self) -> dict[str, PipelineScheduler]: mapping = {} for conf_id, ab_config in self.astrbot_config_mgr.confs.items(): scheduler = PipelineScheduler( - PipelineContext(ab_config, self.plugin_manager, conf_id), + PipelineContext(ab_config, self.plugin_manager, conf_id, self.db), ) await scheduler.initialize() mapping[conf_id] = scheduler @@ -479,7 +479,7 @@ async def reload_pipeline_scheduler(self, conf_id: str) -> None: if not ab_config: raise ValueError(f"配置文件 {conf_id} 不存在") scheduler = PipelineScheduler( - PipelineContext(ab_config, self.plugin_manager, conf_id), + PipelineContext(ab_config, self.plugin_manager, conf_id, self.db), ) await scheduler.initialize() self.pipeline_scheduler_mapping[conf_id] = scheduler diff --git a/astrbot/core/db/__init__.py b/astrbot/core/db/__init__.py index 29053717e0..113218dd70 100644 --- a/astrbot/core/db/__init__.py +++ b/astrbot/core/db/__init__.py @@ -850,6 +850,22 @@ async def upsert_umo_alias( """Create or update the display alias metadata for a UMO.""" ... + @abc.abstractmethod + async def upsert_umo_auto_name( + self, + umo: str, + creator_sender_id: str, + auto_name: str, + ) -> None: + """Create or update only the automatically discovered UMO name. + + Args: + umo: Unified message origin to name. + creator_sender_id: Sender that first caused the UMO to be recorded. + auto_name: Name discovered from the inbound platform message. + """ + ... + @abc.abstractmethod async def get_umo_alias(self, umo: str) -> UmoAlias | None: """Get alias metadata for one UMO.""" diff --git a/astrbot/core/db/po.py b/astrbot/core/db/po.py index 2a4a013806..366a07292d 100644 --- a/astrbot/core/db/po.py +++ b/astrbot/core/db/po.py @@ -346,7 +346,7 @@ class UmoAlias(TimestampMixin, SQLModel, table=True): sa_column_kwargs={"autoincrement": True}, default=None, ) - umo: str = Field(nullable=False, max_length=512, unique=True, index=True) + umo: str = Field(nullable=False, max_length=512) creator_sender_id: str = Field(nullable=False, max_length=255) auto_name: str | None = Field(default=None, max_length=255) user_alias: str | None = Field(default=None, max_length=255) diff --git a/astrbot/core/db/sqlite.py b/astrbot/core/db/sqlite.py index a198cae5e5..c40d710e44 100644 --- a/astrbot/core/db/sqlite.py +++ b/astrbot/core/db/sqlite.py @@ -10,6 +10,7 @@ from deprecated import deprecated from sqlalchemy import CursorResult, Row, case, not_ from sqlalchemy.dialects.sqlite import dialect as sqlite_dialect +from sqlalchemy.dialects.sqlite import insert as sqlite_insert from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import defer from sqlmodel import col, delete, desc, func, or_, select, text, update @@ -72,6 +73,9 @@ async def initialize(self) -> None: await self._ensure_platform_message_history_checkpoint_column(conn) await self._ensure_chatui_project_workspace_columns(conn) await self._ensure_conversation_indexes(conn) + # The table-level unique constraint already provides an index for UMO + # lookups. Older schemas also created this redundant explicit index. + await conn.execute(text("DROP INDEX IF EXISTS ix_umo_aliases_umo")) await conn.commit() async def _ensure_conversation_indexes(self, conn) -> None: @@ -2152,30 +2156,80 @@ async def upsert_umo_alias( auto_name: str | None, user_alias: str | None, ) -> UmoAlias: - """Create or update alias metadata for a UMO.""" + """Create or replace user-controlled alias metadata for a UMO. + + Args: + umo: Unified message origin to name. + creator_sender_id: Sender responsible for the manual alias update. + auto_name: Latest name discovered from platform metadata. + user_alias: User-controlled display alias. + + Returns: + Persisted UMO alias record. + """ + now = datetime.now(timezone.utc) + statement = sqlite_insert(UmoAlias).values( + umo=umo, + creator_sender_id=creator_sender_id, + auto_name=auto_name, + user_alias=user_alias, + created_at=now, + updated_at=now, + ) + statement = statement.on_conflict_do_update( + index_elements=[UmoAlias.umo], + set_={ + "creator_sender_id": statement.excluded.creator_sender_id, + "auto_name": statement.excluded.auto_name, + "user_alias": statement.excluded.user_alias, + "updated_at": now, + }, + ) async with self.get_db() as session: session: AsyncSession async with session.begin(): + await session.execute(statement) result = await session.execute( select(UmoAlias).where(col(UmoAlias.umo) == umo) ) - alias = result.scalar_one_or_none() - if alias: - alias.creator_sender_id = creator_sender_id - alias.auto_name = auto_name - alias.user_alias = user_alias - alias.updated_at = datetime.now(timezone.utc) - else: - alias = UmoAlias( - umo=umo, - creator_sender_id=creator_sender_id, - auto_name=auto_name, - user_alias=user_alias, - ) - session.add(alias) - await session.flush() - await session.refresh(alias) - return alias + return result.scalar_one() + + async def upsert_umo_auto_name( + self, + umo: str, + creator_sender_id: str, + auto_name: str, + ) -> None: + """Persist an automatic UMO name without changing its manual alias. + + Args: + umo: Unified message origin to name. + creator_sender_id: Sender that first caused the UMO to be recorded. + auto_name: Name discovered from the inbound platform message. + """ + now = datetime.now(timezone.utc) + statement = sqlite_insert(UmoAlias).values( + umo=umo, + creator_sender_id=creator_sender_id, + auto_name=auto_name, + user_alias=None, + created_at=now, + updated_at=now, + ) + statement = statement.on_conflict_do_update( + index_elements=[UmoAlias.umo], + set_={ + "auto_name": statement.excluded.auto_name, + "updated_at": now, + }, + where=col(UmoAlias.auto_name).is_distinct_from( + statement.excluded.auto_name + ), + ) + async with self.get_db() as session: + session: AsyncSession + async with session.begin(): + await session.execute(statement) async def get_umo_alias(self, umo: str) -> UmoAlias | None: """Get alias metadata for one UMO.""" diff --git a/astrbot/core/pipeline/context.py b/astrbot/core/pipeline/context.py index 47cd33b238..1734d5a2fd 100644 --- a/astrbot/core/pipeline/context.py +++ b/astrbot/core/pipeline/context.py @@ -8,6 +8,7 @@ from .context_utils import call_event_hook, call_handler if TYPE_CHECKING: + from astrbot.core.db import BaseDatabase from astrbot.core.star import PluginManager @@ -18,5 +19,6 @@ class PipelineContext: astrbot_config: AstrBotConfig # AstrBot 配置对象 plugin_manager: PluginManager # 插件管理器对象 astrbot_config_id: str + db_helper: BaseDatabase | None = None call_handler = call_handler call_event_hook = call_event_hook diff --git a/astrbot/core/pipeline/waking_check/stage.py b/astrbot/core/pipeline/waking_check/stage.py index 6c09efb5de..f02915bd64 100644 --- a/astrbot/core/pipeline/waking_check/stage.py +++ b/astrbot/core/pipeline/waking_check/stage.py @@ -13,6 +13,7 @@ from ..context import PipelineContext from ..stage import Stage, register_stage +from .umo_auto_name import UmoAutoNameRecorder UNIQUE_SESSION_ID_BUILDERS: dict[str, Callable[[AstrMessageEvent], str | None]] = { "aiocqhttp": lambda e: f"{e.get_sender_id()}_{e.get_group_id()}", @@ -73,6 +74,10 @@ async def initialize(self, ctx: PipelineContext) -> None: ) platform_settings = self.ctx.astrbot_config.get("platform_settings", {}) self.unique_session = platform_settings.get("unique_session", False) + self._umo_auto_name_recorder = UmoAutoNameRecorder( + ctx.db_helper, + ctx.astrbot_config_id, + ) async def process( self, @@ -218,6 +223,8 @@ async def process( f"{star_map[handler.handler_module_path].name}.", ) event.stop_event() + if event.is_wake: + self._umo_auto_name_recorder.schedule(event) return is_wake = True @@ -244,5 +251,7 @@ async def process( event.set_extra("activated_handlers", activated_handlers) event.set_extra("handlers_parsed_params", handlers_parsed_params) - if not is_wake: + if is_wake: + self._umo_auto_name_recorder.schedule(event) + else: event.stop_event() diff --git a/astrbot/core/pipeline/waking_check/umo_auto_name.py b/astrbot/core/pipeline/waking_check/umo_auto_name.py new file mode 100644 index 0000000000..03dc444547 --- /dev/null +++ b/astrbot/core/pipeline/waking_check/umo_auto_name.py @@ -0,0 +1,110 @@ +from __future__ import annotations + +import asyncio +from collections import OrderedDict +from typing import TYPE_CHECKING + +from astrbot import logger +from astrbot.core.umo_alias import get_event_auto_name + +if TYPE_CHECKING: + from astrbot.core.db import BaseDatabase + from astrbot.core.platform.astr_message_event import AstrMessageEvent + +MAX_UMO_AUTO_NAME_CACHE_SIZE = 10_000 + + +class UmoAutoNameRecorder: + """Persist changed UMO names without blocking the waking stage.""" + + def __init__( + self, + db_helper: BaseDatabase | None, + config_id: str, + ) -> None: + """Initialize the bounded cache and background writer state. + + Args: + db_helper: Database used to persist automatic names. + config_id: Pipeline configuration identifier used in the task name. + """ + self.db_helper = db_helper + self.config_id = config_id + self._cache: OrderedDict[str, str] = OrderedDict() + self._pending: OrderedDict[str, tuple[str, str]] = OrderedDict() + self._writer_task: asyncio.Task[None] | None = None + + def schedule(self, event: AstrMessageEvent) -> None: + """Queue a changed automatic name from an awakened event. + + Args: + event: Awakened event containing the UMO and display metadata. + """ + if self.db_helper is None: + return + + umo = event.unified_msg_origin + auto_name = get_event_auto_name(event, fallback_to_id=False) + if not auto_name: + return + if self._cache.get(umo) == auto_name: + self._cache.move_to_end(umo) + return + + self._cache[umo] = auto_name + self._cache.move_to_end(umo) + if len(self._cache) > MAX_UMO_AUTO_NAME_CACHE_SIZE: + self._cache.popitem(last=False) + + self._pending[umo] = (str(event.get_sender_id() or ""), auto_name) + self._pending.move_to_end(umo) + if len(self._pending) > MAX_UMO_AUTO_NAME_CACHE_SIZE: + dropped_umo, (_, dropped_name) = self._pending.popitem(last=False) + if self._cache.get(dropped_umo) == dropped_name: + self._cache.pop(dropped_umo, None) + + if self._writer_task is None or self._writer_task.done(): + task = asyncio.create_task( + self._flush(), + name=f"umo_auto_name_writer:{self.config_id}", + ) + self._writer_task = task + task.add_done_callback(self._on_writer_done) + + async def _flush(self) -> None: + """Persist queued names sequentially, coalescing changes per UMO.""" + if self.db_helper is None: + return + + try: + while self._pending: + umo, (creator_sender_id, auto_name) = self._pending.popitem(last=False) + try: + await self.db_helper.upsert_umo_auto_name( + umo=umo, + creator_sender_id=creator_sender_id, + auto_name=auto_name, + ) + except Exception as exc: + logger.warning( + "Failed to persist automatic UMO name for %s: %s", + umo, + exc, + ) + if umo not in self._pending and self._cache.get(umo) == auto_name: + self._cache.pop(umo, None) + finally: + self._writer_task = None + + @staticmethod + def _on_writer_done(task: asyncio.Task[None]) -> None: + """Expose unexpected writer failures. + + Args: + task: Completed automatic-name writer task. + """ + if task.cancelled(): + return + exc = task.exception() + if exc is not None: + logger.error("UMO automatic-name writer failed.", exc_info=exc) diff --git a/astrbot/core/platform/sources/discord/discord_platform_adapter.py b/astrbot/core/platform/sources/discord/discord_platform_adapter.py index f8205017ee..65c3b3521a 100644 --- a/astrbot/core/platform/sources/discord/discord_platform_adapter.py +++ b/astrbot/core/platform/sources/discord/discord_platform_adapter.py @@ -83,14 +83,11 @@ async def send_by_session( if channel: message_obj.type = self._get_message_type(channel) - message_obj.group_id = self._get_channel_id(channel) - group_name = self._get_group_name(channel) - if ( - message_obj.type == MessageType.GROUP_MESSAGE - and message_obj.group - and group_name - ): - message_obj.group.group_name = group_name + if message_obj.type == MessageType.GROUP_MESSAGE: + message_obj.group_id = self._get_channel_id(channel) + group_name = self._get_group_name(channel) + if message_obj.group and group_name: + message_obj.group.group_name = group_name else: logger.warning( f"[Discord] Can't get channel info for {channel_id_str}, will guess message type.", @@ -253,10 +250,11 @@ def _convert_message_to_abm(self, data: dict) -> AstrBotMessage: abm = AstrBotMessage() abm.type = self._get_message_type(message.channel) - abm.group_id = self._get_channel_id(message.channel) - group_name = self._get_group_name(message.channel) - if abm.type == MessageType.GROUP_MESSAGE and abm.group and group_name: - abm.group.group_name = group_name + if abm.type == MessageType.GROUP_MESSAGE: + abm.group_id = self._get_channel_id(message.channel) + group_name = self._get_group_name(message.channel) + if abm.group and group_name: + abm.group.group_name = group_name abm.message_str = content abm.sender = MessageMember( user_id=str(message.author.id), @@ -542,10 +540,11 @@ async def dynamic_callback( abm = AstrBotMessage() if channel is not None: abm.type = self._get_message_type(channel, ctx.guild_id) - abm.group_id = self._get_channel_id(channel) - group_name = self._get_group_name(channel) - if abm.type == MessageType.GROUP_MESSAGE and abm.group and group_name: - abm.group.group_name = group_name + if abm.type == MessageType.GROUP_MESSAGE: + abm.group_id = self._get_channel_id(channel) + group_name = self._get_group_name(channel) + if abm.group and group_name: + abm.group.group_name = group_name else: # 防守式兜底:channel 取不到时,仍能根据 guild_id/channel_id 推断会话信息 abm.type = ( @@ -553,7 +552,8 @@ async def dynamic_callback( if ctx.guild_id is not None else MessageType.FRIEND_MESSAGE ) - abm.group_id = str(ctx.channel_id) + if abm.type == MessageType.GROUP_MESSAGE: + abm.group_id = str(ctx.channel_id) abm.message_str = message_str_for_filter abm.sender = MessageMember( diff --git a/astrbot/core/umo_alias.py b/astrbot/core/umo_alias.py index 8d8f01ee63..7ef394d528 100644 --- a/astrbot/core/umo_alias.py +++ b/astrbot/core/umo_alias.py @@ -21,20 +21,34 @@ def parse_umo(umo: Any) -> dict[str, str]: } -def get_event_auto_name(event: Any) -> str: +def get_event_auto_name(event: Any, *, fallback_to_id: bool = True) -> str: + """Resolve an automatic display name from inbound event metadata. + + Args: + event: Platform event containing group and sender metadata. + fallback_to_id: Whether to use the group or sender ID when no name exists. + + Returns: + Normalized group or sender name, an optional ID fallback, or an empty string. + """ group_id = event.get_group_id() if hasattr(event, "get_group_id") else "" message_obj = getattr(event, "message_obj", None) group = getattr(message_obj, "group", None) if group_id: group_name = normalize_umo_name(getattr(group, "group_name", None)) - return group_name or normalize_umo_name(group_id) + if group_name: + return group_name + return normalize_umo_name(group_id) if fallback_to_id else "" sender_name = "" if hasattr(event, "get_sender_name"): sender_name = event.get_sender_name() sender_id = event.get_sender_id() if hasattr(event, "get_sender_id") else "" - return normalize_umo_name(sender_name) or normalize_umo_name(sender_id) + normalized_sender_name = normalize_umo_name(sender_name) + if normalized_sender_name: + return normalized_sender_name + return normalize_umo_name(sender_id) if fallback_to_id else "" def get_umo_display_name( diff --git a/tests/test_discord_adapter.py b/tests/test_discord_adapter.py index 0d4bdfd66f..b0413965da 100644 --- a/tests/test_discord_adapter.py +++ b/tests/test_discord_adapter.py @@ -67,8 +67,9 @@ async def test_discord_private_message_does_not_get_group_name(): abm = await adapter.convert_message({"message": message}) assert abm.type == MessageType.FRIEND_MESSAGE - assert abm.group is not None - assert abm.group.group_name is None + assert abm.group is None + assert abm.group_id == "" + assert abm.sender.nickname == "tester" def test_discord_group_name_falls_back_when_one_name_is_missing(): diff --git a/tests/test_umo_alias.py b/tests/test_umo_alias.py index 1cbeaf9468..5ebd55b52e 100644 --- a/tests/test_umo_alias.py +++ b/tests/test_umo_alias.py @@ -1,9 +1,11 @@ +import asyncio import importlib import sys from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from sqlmodel import text from astrbot.builtin_stars.builtin_commands.commands.name import NameCommand from astrbot.core.star.filter.permission import PermissionType, PermissionTypeFilter @@ -58,6 +60,99 @@ async def test_umo_alias_upsert_updates_existing_record(temp_db): assert serialize_umo_alias(fetched, fetched.umo)["display_name"] == "New Alias" +@pytest.mark.asyncio +async def test_auto_name_upsert_preserves_manual_alias_and_skips_unchanged_update( + temp_db, +): + await temp_db.upsert_umo_alias( + umo="qq:GroupMessage:1000", + creator_sender_id="admin-1", + auto_name="Engineering Group", + user_alias="Backend Room", + ) + before = await temp_db.get_umo_alias("qq:GroupMessage:1000") + assert before is not None + + await temp_db.upsert_umo_auto_name( + umo="qq:GroupMessage:1000", + creator_sender_id="sender-2", + auto_name="Engineering Group", + ) + unchanged = await temp_db.get_umo_alias("qq:GroupMessage:1000") + assert unchanged is not None + assert unchanged.updated_at == before.updated_at + + await temp_db.upsert_umo_auto_name( + umo="qq:GroupMessage:1000", + creator_sender_id="sender-3", + auto_name="Renamed Group", + ) + renamed = await temp_db.get_umo_alias("qq:GroupMessage:1000") + assert renamed is not None + assert renamed.auto_name == "Renamed Group" + assert renamed.user_alias == "Backend Room" + assert renamed.creator_sender_id == "admin-1" + assert renamed.updated_at != before.updated_at + + +@pytest.mark.asyncio +async def test_auto_name_upsert_creates_alias_without_manual_name(temp_db): + await temp_db.upsert_umo_auto_name( + umo="qq:FriendMessage:sender-1", + creator_sender_id="sender-1", + auto_name="Alice", + ) + + alias = await temp_db.get_umo_alias("qq:FriendMessage:sender-1") + assert alias is not None + assert alias.creator_sender_id == "sender-1" + assert alias.auto_name == "Alice" + assert alias.user_alias is None + + +@pytest.mark.asyncio +async def test_concurrent_manual_and_auto_name_upserts_preserve_manual_alias(temp_db): + assert await temp_db.get_umo_alias("qq:GroupMessage:1000") is None + + await asyncio.gather( + temp_db.upsert_umo_alias( + umo="qq:GroupMessage:1000", + creator_sender_id="admin-1", + auto_name="Engineering Group", + user_alias="Backend Room", + ), + temp_db.upsert_umo_auto_name( + umo="qq:GroupMessage:1000", + creator_sender_id="sender-1", + auto_name="Engineering Group", + ), + ) + + alias = await temp_db.get_umo_alias("qq:GroupMessage:1000") + assert alias is not None + assert alias.creator_sender_id == "admin-1" + assert alias.auto_name == "Engineering Group" + assert alias.user_alias == "Backend Room" + + +@pytest.mark.asyncio +async def test_database_initialization_removes_redundant_umo_index(temp_db): + async with temp_db.get_db() as session: + async with session.begin(): + await session.execute( + text("CREATE UNIQUE INDEX ix_umo_aliases_umo ON umo_aliases (umo)") + ) + + await temp_db.initialize() + + async with temp_db.get_db() as session: + result = await session.execute(text("PRAGMA index_list(umo_aliases)")) + indexes = result.fetchall() + + assert "ix_umo_aliases_umo" not in {row[1] for row in indexes} + assert any(row[2] for row in indexes) + + @pytest.mark.asyncio async def test_name_command_saves_group_alias_with_auto_name(temp_db): context = SimpleNamespace(get_db=lambda: temp_db) @@ -159,6 +254,24 @@ def test_umo_name_helpers_accept_numeric_ids(): ) +def test_event_auto_name_can_skip_group_and_sender_id_fallbacks(): + group_event = SimpleNamespace( + message_obj=SimpleNamespace(group=SimpleNamespace(group_name=None)), + get_group_id=lambda: "group-1", + get_sender_id=lambda: "sender-1", + get_sender_name=lambda: "Alice", + ) + friend_event = SimpleNamespace( + message_obj=SimpleNamespace(group=None), + get_group_id=lambda: "", + get_sender_id=lambda: "sender-1", + get_sender_name=lambda: "", + ) + + assert get_event_auto_name(group_event, fallback_to_id=False) == "" + assert get_event_auto_name(friend_event, fallback_to_id=False) == "" + + def test_parse_umo_handles_empty_values(): assert parse_umo(None) == { "platform": "unknown", diff --git a/tests/unit/test_waking_check_api_key_admin.py b/tests/unit/test_waking_check_api_key_admin.py index 0d85ce22fe..ee009c1e97 100644 --- a/tests/unit/test_waking_check_api_key_admin.py +++ b/tests/unit/test_waking_check_api_key_admin.py @@ -39,6 +39,7 @@ async def test_waking_check_enforces_api_key_admin_authorization( stage.ignore_at_all = False stage.disable_builtin_commands = False stage.no_permission_reply = True + stage._umo_auto_name_recorder = MagicMock() event = MagicMock() event.message_str = "hello" @@ -48,9 +49,7 @@ async def test_waking_check_enforces_api_key_admin_authorization( event.is_private_chat.return_value = True event.get_platform_name.return_value = "webchat" event.get_extra.side_effect = lambda key=None, default=None: ( - api_key_allow_admin_role - if key == "_api_key_allow_admin_role" - else default + api_key_allow_admin_role if key == "_api_key_allow_admin_role" else default ) monkeypatch.setattr( star_handlers_registry, diff --git a/tests/unit/test_waking_check_umo_alias.py b/tests/unit/test_waking_check_umo_alias.py new file mode 100644 index 0000000000..b74f9cc2cf --- /dev/null +++ b/tests/unit/test_waking_check_umo_alias.py @@ -0,0 +1,222 @@ +"""Tests for automatic UMO names recorded by the waking stage.""" + +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from astrbot.core.pipeline.waking_check.stage import WakingCheckStage +from astrbot.core.pipeline.waking_check.umo_auto_name import UmoAutoNameRecorder +from astrbot.core.platform.message_type import MessageType +from astrbot.core.star.session_plugin_manager import SessionPluginManager + + +def make_group_event(group_id: str, group_name: str | None, message: str = "/hello"): + """Create a group event carrying wake and display metadata. + + Args: + group_id: Platform group identifier. + group_name: Platform group display name. + message: Event message text. + + Returns: + Mocked group message event. + """ + event = MagicMock() + event.unified_msg_origin = f"test-platform:GroupMessage:{group_id}" + event.message_obj = SimpleNamespace( + type=MessageType.GROUP_MESSAGE, + group=SimpleNamespace(group_name=group_name), + ) + event.message_str = message + event.is_wake = False + event.role = "member" + event.get_group_id.return_value = group_id + event.get_sender_id.return_value = "sender-1" + event.get_self_id.return_value = "bot-1" + event.get_messages.return_value = [MagicMock()] + event.is_private_chat.return_value = False + event.get_platform_name.return_value = "test-platform" + event.get_extra.side_effect = lambda key=None, default=None: default + return event + + +async def make_stage(db_helper: MagicMock) -> WakingCheckStage: + """Initialize a waking stage with automatic-name persistence enabled. + + Args: + db_helper: Mock database used by the stage writer. + + Returns: + Initialized waking stage. + """ + stage = WakingCheckStage() + await stage.initialize( + SimpleNamespace( + astrbot_config={ + "admins_id": [], + "wake_prefix": ["/"], + "plugin_set": ["*"], + "platform_settings": { + "friend_message_needs_wake_prefix": True, + }, + }, + astrbot_config_id="test-conf-id", + db_helper=db_helper, + ) + ) + return stage + + +@pytest.mark.asyncio +async def test_waking_stage_records_only_awakened_events(monkeypatch): + """Record a name immediately after waking and ignore ambient messages.""" + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock() + stage = await make_stage(db_helper) + monkeypatch.setattr( + "astrbot.core.pipeline.waking_check.stage.star_handlers_registry.get_handlers_by_event_type", + lambda *_args, **_kwargs: [], + ) + + async def return_handlers(_event, handlers): + return handlers + + monkeypatch.setattr( + SessionPluginManager, + "filter_handlers_by_session", + return_handlers, + ) + + ignored_event = make_group_event("group-1", "Engineering", "hello") + await stage.process(ignored_event) + assert stage._umo_auto_name_recorder._writer_task is None + + awakened_event = make_group_event("group-1", "Engineering") + await stage.process(awakened_event) + writer_task = stage._umo_auto_name_recorder._writer_task + assert writer_task is not None + await writer_task + + db_helper.upsert_umo_auto_name.assert_awaited_once_with( + umo="test-platform:GroupMessage:group-1", + creator_sender_id="sender-1", + auto_name="Engineering", + ) + + +@pytest.mark.asyncio +async def test_waking_stage_coalesces_auto_name_changes(): + """Persist only the latest name from an event burst for one UMO.""" + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock() + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") + + for group_name in ("Engineering", "Engineering", "Renamed"): + recorder.schedule(make_group_event("group-1", group_name)) + + writer_task = recorder._writer_task + assert writer_task is not None + await writer_task + + db_helper.upsert_umo_auto_name.assert_awaited_once_with( + umo="test-platform:GroupMessage:group-1", + creator_sender_id="sender-1", + auto_name="Renamed", + ) + + +@pytest.mark.asyncio +async def test_waking_stage_skips_missing_group_and_sender_names(): + """Do not persist ID fallbacks when platform names are unavailable.""" + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock() + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") + + recorder.schedule(make_group_event("group-1", None)) + + friend_event = MagicMock() + friend_event.unified_msg_origin = "test-platform:FriendMessage:sender-2" + friend_event.message_obj = SimpleNamespace(group=None) + friend_event.get_group_id.return_value = "" + friend_event.get_sender_name.return_value = "" + friend_event.get_sender_id.return_value = "sender-2" + recorder.schedule(friend_event) + + assert recorder._writer_task is None + assert not recorder._cache + db_helper.upsert_umo_auto_name.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_waking_stage_bounds_auto_name_cache(): + """Evict old UMO names when the per-stage cache reaches its bound.""" + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock() + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") + + with patch( + "astrbot.core.pipeline.waking_check.umo_auto_name.MAX_UMO_AUTO_NAME_CACHE_SIZE", + 2, + ): + for index in range(3): + recorder.schedule(make_group_event(f"group-{index}", f"Group {index}")) + + writer_task = recorder._writer_task + assert writer_task is not None + await writer_task + + assert list(recorder._cache) == [ + "test-platform:GroupMessage:group-1", + "test-platform:GroupMessage:group-2", + ] + assert db_helper.upsert_umo_auto_name.await_count == 2 + + +@pytest.mark.asyncio +async def test_waking_stage_retries_after_database_failure(): + """Evict a failed cache entry so a later wake retries the write.""" + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock( + side_effect=[RuntimeError("database unavailable"), None] + ) + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") + event = make_group_event("group-1", "Engineering") + + with patch("astrbot.core.pipeline.waking_check.umo_auto_name.logger"): + recorder.schedule(event) + first_writer = recorder._writer_task + assert first_writer is not None + await first_writer + + assert event.unified_msg_origin not in recorder._cache + + recorder.schedule(event) + second_writer = recorder._writer_task + assert second_writer is not None + await second_writer + + assert db_helper.upsert_umo_auto_name.await_count == 2 + + +@pytest.mark.asyncio +async def test_waking_stage_writer_does_not_block_processing(): + """Return from the waking stage while its database writer is blocked.""" + database_started = asyncio.Event() + release_database = asyncio.Event() + + async def block_database_write(**kwargs): # noqa: ARG001 + database_started.set() + await release_database.wait() + + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock(side_effect=block_database_write) + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") + recorder.schedule(make_group_event("group-1", "Engineering")) + + await asyncio.wait_for(database_started.wait(), timeout=1.0) + release_database.set() + writer_task = recorder._writer_task + if writer_task is not None: + await writer_task