From 1a93bff49d5affbd36c33a191d46701eb7b700bc Mon Sep 17 00:00:00 2001 From: Soulter <905617992@qq.com> Date: Tue, 1 Sep 2026 22:27:37 +0800 Subject: [PATCH 1/4] feat: automatically record UMO names --- astrbot/core/core_lifecycle.py | 1 + astrbot/core/db/__init__.py | 16 + astrbot/core/db/po.py | 2 +- astrbot/core/db/sqlite.py | 90 ++++-- astrbot/core/event_bus.py | 87 ++++++ .../discord/discord_platform_adapter.py | 34 +-- astrbot/core/umo_alias.py | 20 +- tests/test_discord_adapter.py | 5 +- tests/test_umo_alias.py | 113 +++++++ tests/unit/test_event_bus.py | 279 ++++++++++++++++++ 10 files changed, 606 insertions(+), 41 deletions(-) diff --git a/astrbot/core/core_lifecycle.py b/astrbot/core/core_lifecycle.py index db8a6ddc1b..a632964792 100644 --- a/astrbot/core/core_lifecycle.py +++ b/astrbot/core/core_lifecycle.py @@ -278,6 +278,7 @@ async def initialize(self) -> None: self.event_queue, self.pipeline_scheduler_mapping, self.astrbot_config_mgr, + self.db, ) # 记录启动时间 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/event_bus.py b/astrbot/core/event_bus.py index aa7b80f937..34706dcde9 100644 --- a/astrbot/core/event_bus.py +++ b/astrbot/core/event_bus.py @@ -12,13 +12,18 @@ import asyncio from asyncio import Queue +from collections import OrderedDict from astrbot.core import logger from astrbot.core.astrbot_config_mgr import AstrBotConfigManager +from astrbot.core.db import BaseDatabase from astrbot.core.pipeline.scheduler import PipelineScheduler +from astrbot.core.umo_alias import get_event_auto_name from .platform import AstrMessageEvent +MAX_UMO_AUTO_NAME_CACHE_SIZE = 10_000 + class EventBus: """用于处理事件的分发和处理""" @@ -28,11 +33,16 @@ def __init__( event_queue: Queue, pipeline_scheduler_mapping: dict[str, PipelineScheduler], astrbot_config_mgr: AstrBotConfigManager, + db_helper: BaseDatabase | None = None, ) -> None: self.event_queue = event_queue # 事件队列 # abconf uuid -> scheduler self.pipeline_scheduler_mapping = pipeline_scheduler_mapping self.astrbot_config_mgr = astrbot_config_mgr + self.db_helper = db_helper + self._umo_auto_name_cache: OrderedDict[str, str] = OrderedDict() + self._pending_umo_auto_names: OrderedDict[str, tuple[str, str]] = OrderedDict() + self._umo_auto_name_writer_task: asyncio.Task[None] | None = None # 持有正在执行的 pipeline 任务的强引用, 防止 task 在 pending 状态被 GC 回收 self._pending_tasks: set[asyncio.Task] = set() @@ -49,10 +59,87 @@ async def dispatch(self) -> None: f"PipelineScheduler not found for id: {conf_id}, event ignored." ) continue + self._schedule_umo_auto_name_recording(event) task = asyncio.create_task(scheduler.execute(event)) self._pending_tasks.add(task) task.add_done_callback(self._on_task_done) + def _schedule_umo_auto_name_recording(self, event: AstrMessageEvent) -> None: + """Queue a changed automatic UMO name for background persistence. + + Args: + event: Inbound platform 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._umo_auto_name_cache.get(umo) == auto_name: + self._umo_auto_name_cache.move_to_end(umo) + return + + self._umo_auto_name_cache[umo] = auto_name + self._umo_auto_name_cache.move_to_end(umo) + if len(self._umo_auto_name_cache) > MAX_UMO_AUTO_NAME_CACHE_SIZE: + self._umo_auto_name_cache.popitem(last=False) + + self._pending_umo_auto_names[umo] = ( + str(event.get_sender_id() or ""), + auto_name, + ) + self._pending_umo_auto_names.move_to_end(umo) + if len(self._pending_umo_auto_names) > MAX_UMO_AUTO_NAME_CACHE_SIZE: + dropped_umo, (_, dropped_name) = self._pending_umo_auto_names.popitem( + last=False + ) + if self._umo_auto_name_cache.get(dropped_umo) == dropped_name: + self._umo_auto_name_cache.pop(dropped_umo, None) + + if ( + self._umo_auto_name_writer_task is None + or self._umo_auto_name_writer_task.done() + ): + task = asyncio.create_task( + self._flush_umo_auto_names(), + name="umo_auto_name_writer", + ) + self._umo_auto_name_writer_task = task + self._pending_tasks.add(task) + task.add_done_callback(self._on_task_done) + + async def _flush_umo_auto_names(self) -> None: + """Persist queued UMO names sequentially, coalescing changes per UMO.""" + if self.db_helper is None: + return + + try: + while self._pending_umo_auto_names: + umo, (creator_sender_id, auto_name) = ( + self._pending_umo_auto_names.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_umo_auto_names + and self._umo_auto_name_cache.get(umo) == auto_name + ): + self._umo_auto_name_cache.pop(umo, None) + finally: + self._umo_auto_name_writer_task = None + def _on_task_done(self, task: asyncio.Task) -> None: """pipeline 任务结束回调: 移除强引用并暴露未捕获的异常""" self._pending_tasks.discard(task) 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_event_bus.py b/tests/unit/test_event_bus.py index 1ecdbf1e31..8ccaa66135 100644 --- a/tests/unit/test_event_bus.py +++ b/tests/unit/test_event_bus.py @@ -2,6 +2,7 @@ import asyncio from contextlib import suppress +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -101,6 +102,284 @@ async def execute_and_signal(event): # noqa: ARG001 "test-platform:group:123" ) + @pytest.mark.asyncio + async def test_dispatch_coalesces_auto_name_changes( + self, + event_queue, + mock_pipeline_scheduler, + mock_config_manager, + ): + """Persist only the latest automatic name from a queued event burst.""" + processed = asyncio.Event() + processed_count = 0 + + async def execute_and_count(event): # noqa: ARG001 + nonlocal processed_count + processed_count += 1 + if processed_count == 3: + processed.set() + + mock_pipeline_scheduler.execute.side_effect = execute_and_count + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock() + bus = EventBus( + event_queue=event_queue, + pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, + astrbot_config_mgr=mock_config_manager, + db_helper=db_helper, + ) + + for group_name in ("Engineering Group", "Engineering Group", "Renamed"): + mock_event = MagicMock() + mock_event.unified_msg_origin = "test-platform:GroupMessage:group-1" + mock_event.message_obj = SimpleNamespace( + group=SimpleNamespace(group_name=group_name) + ) + mock_event.get_group_id.return_value = "group-1" + mock_event.get_platform_id.return_value = "test-platform" + mock_event.get_platform_name.return_value = "Test Platform" + mock_event.get_sender_name.return_value = "Alice" + mock_event.get_sender_id.return_value = "sender-1" + mock_event.get_message_outline.return_value = "Hello" + await event_queue.put(mock_event) + + task = asyncio.create_task(bus.dispatch()) + try: + await asyncio.wait_for(processed.wait(), timeout=1.0) + finally: + task.cancel() + with suppress(asyncio.CancelledError): + await task + + writer_task = bus._umo_auto_name_writer_task + if writer_task is not None: + await writer_task + assert db_helper.upsert_umo_auto_name.await_count == 1 + assert [ + call.kwargs["auto_name"] + for call in db_helper.upsert_umo_auto_name.await_args_list + ] == ["Renamed"] + + @pytest.mark.asyncio + async def test_dispatch_does_not_wait_for_auto_name_database_write( + self, + event_queue, + mock_pipeline_scheduler, + mock_config_manager, + ): + """Start the pipeline while the background alias writer is blocked.""" + database_started = asyncio.Event() + release_database = asyncio.Event() + pipeline_started = asyncio.Event() + + async def block_database_write(**kwargs): # noqa: ARG001 + database_started.set() + await release_database.wait() + + async def execute_and_signal(event): # noqa: ARG001 + pipeline_started.set() + + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock(side_effect=block_database_write) + mock_pipeline_scheduler.execute.side_effect = execute_and_signal + bus = EventBus( + event_queue=event_queue, + pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, + astrbot_config_mgr=mock_config_manager, + db_helper=db_helper, + ) + mock_event = MagicMock() + mock_event.unified_msg_origin = "test-platform:GroupMessage:group-1" + mock_event.message_obj = SimpleNamespace( + group=SimpleNamespace(group_name="Engineering Group") + ) + mock_event.get_group_id.return_value = "group-1" + mock_event.get_platform_id.return_value = "test-platform" + mock_event.get_platform_name.return_value = "Test Platform" + mock_event.get_sender_name.return_value = "Alice" + mock_event.get_sender_id.return_value = "sender-1" + mock_event.get_message_outline.return_value = "Hello" + await event_queue.put(mock_event) + + task = asyncio.create_task(bus.dispatch()) + try: + await asyncio.wait_for(database_started.wait(), timeout=1.0) + await asyncio.wait_for(pipeline_started.wait(), timeout=1.0) + finally: + release_database.set() + writer_task = bus._umo_auto_name_writer_task + if writer_task is not None: + await writer_task + task.cancel() + with suppress(asyncio.CancelledError): + await task + + @pytest.mark.asyncio + async def test_dispatch_bounds_auto_name_cache( + self, + event_queue, + mock_pipeline_scheduler, + mock_config_manager, + ): + """Evict the least recently used UMO when the cache reaches its bound.""" + processed = asyncio.Event() + processed_count = 0 + + async def execute_and_count(event): # noqa: ARG001 + nonlocal processed_count + processed_count += 1 + if processed_count == 3: + processed.set() + + mock_pipeline_scheduler.execute.side_effect = execute_and_count + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock() + with patch("astrbot.core.event_bus.MAX_UMO_AUTO_NAME_CACHE_SIZE", 2): + bus = EventBus( + event_queue=event_queue, + pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, + astrbot_config_mgr=mock_config_manager, + db_helper=db_helper, + ) + + for index in range(3): + mock_event = MagicMock() + mock_event.unified_msg_origin = ( + f"test-platform:GroupMessage:group-{index}" + ) + mock_event.message_obj = SimpleNamespace( + group=SimpleNamespace(group_name=f"Group {index}") + ) + mock_event.get_group_id.return_value = f"group-{index}" + mock_event.get_platform_id.return_value = "test-platform" + mock_event.get_platform_name.return_value = "Test Platform" + mock_event.get_sender_name.return_value = "Alice" + mock_event.get_sender_id.return_value = "sender-1" + mock_event.get_message_outline.return_value = "Hello" + await event_queue.put(mock_event) + + task = asyncio.create_task(bus.dispatch()) + try: + await asyncio.wait_for(processed.wait(), timeout=1.0) + finally: + task.cancel() + with suppress(asyncio.CancelledError): + await task + + assert list(bus._umo_auto_name_cache) == [ + "test-platform:GroupMessage:group-1", + "test-platform:GroupMessage:group-2", + ] + assert not bus._pending_umo_auto_names + assert db_helper.upsert_umo_auto_name.await_count == 2 + + @pytest.mark.asyncio + async def test_auto_name_writer_retries_after_database_failure( + self, + event_queue, + mock_pipeline_scheduler, + mock_config_manager, + ): + """Evict a failed optimistic cache entry so the next event retries.""" + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock( + side_effect=[RuntimeError("database unavailable"), None] + ) + bus = EventBus( + event_queue=event_queue, + pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, + astrbot_config_mgr=mock_config_manager, + db_helper=db_helper, + ) + mock_event = MagicMock() + mock_event.unified_msg_origin = "test-platform:GroupMessage:group-1" + mock_event.message_obj = SimpleNamespace( + group=SimpleNamespace(group_name="Engineering Group") + ) + mock_event.get_group_id.return_value = "group-1" + mock_event.get_sender_id.return_value = "sender-1" + + with patch("astrbot.core.event_bus.logger"): + bus._schedule_umo_auto_name_recording(mock_event) + first_writer = bus._umo_auto_name_writer_task + assert first_writer is not None + await first_writer + + assert mock_event.unified_msg_origin not in bus._umo_auto_name_cache + + bus._schedule_umo_auto_name_recording(mock_event) + second_writer = bus._umo_auto_name_writer_task + assert second_writer is not None + await second_writer + + assert db_helper.upsert_umo_auto_name.await_count == 2 + assert ( + bus._umo_auto_name_cache[mock_event.unified_msg_origin] + == "Engineering Group" + ) + + @pytest.mark.asyncio + async def test_dispatch_skips_missing_group_and_sender_names( + self, + event_queue, + mock_pipeline_scheduler, + mock_config_manager, + ): + """Do not persist ID fallbacks as automatic UMO names.""" + processed = asyncio.Event() + processed_count = 0 + + async def execute_and_count(event): # noqa: ARG001 + nonlocal processed_count + processed_count += 1 + if processed_count == 2: + processed.set() + + mock_pipeline_scheduler.execute.side_effect = execute_and_count + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock() + bus = EventBus( + event_queue=event_queue, + pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, + astrbot_config_mgr=mock_config_manager, + db_helper=db_helper, + ) + + group_event = MagicMock() + group_event.unified_msg_origin = "test-platform:GroupMessage:group-1" + group_event.message_obj = SimpleNamespace( + group=SimpleNamespace(group_name=None) + ) + group_event.get_group_id.return_value = "group-1" + group_event.get_sender_name.return_value = "Alice" + group_event.get_sender_id.return_value = "sender-1" + group_event.get_platform_id.return_value = "test-platform" + group_event.get_platform_name.return_value = "Test Platform" + group_event.get_message_outline.return_value = "Hello" + + 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" + friend_event.get_platform_id.return_value = "test-platform" + friend_event.get_platform_name.return_value = "Test Platform" + friend_event.get_message_outline.return_value = "Hello" + + await event_queue.put(group_event) + await event_queue.put(friend_event) + task = asyncio.create_task(bus.dispatch()) + try: + await asyncio.wait_for(processed.wait(), timeout=1.0) + finally: + task.cancel() + with suppress(asyncio.CancelledError): + await task + + db_helper.upsert_umo_auto_name.assert_not_awaited() + assert not bus._umo_auto_name_cache + @pytest.mark.asyncio async def test_dispatch_handles_missing_scheduler( self, From 5be3fc7f08cf5611a45951a0a65208057013369e Mon Sep 17 00:00:00 2001 From: Soulter <905617992@qq.com> Date: Tue, 1 Sep 2026 22:42:33 +0800 Subject: [PATCH 2/4] refactor: record UMO names only after wake --- astrbot/core/event_bus.py | 18 +++++++++-- tests/unit/test_event_bus.py | 59 +++++++++++++++++++++++++++++++++--- 2 files changed, 71 insertions(+), 6 deletions(-) diff --git a/astrbot/core/event_bus.py b/astrbot/core/event_bus.py index 34706dcde9..043c956833 100644 --- a/astrbot/core/event_bus.py +++ b/astrbot/core/event_bus.py @@ -59,11 +59,25 @@ async def dispatch(self) -> None: f"PipelineScheduler not found for id: {conf_id}, event ignored." ) continue - self._schedule_umo_auto_name_recording(event) - task = asyncio.create_task(scheduler.execute(event)) + task = asyncio.create_task(self._execute_pipeline(scheduler, event)) self._pending_tasks.add(task) task.add_done_callback(self._on_task_done) + async def _execute_pipeline( + self, + scheduler: PipelineScheduler, + event: AstrMessageEvent, + ) -> None: + """Execute the pipeline and record the UMO name after a successful wake. + + Args: + scheduler: Pipeline scheduler selected for the event configuration. + event: Inbound platform event to process. + """ + await scheduler.execute(event) + if event.is_wake: + self._schedule_umo_auto_name_recording(event) + def _schedule_umo_auto_name_recording(self, event: AstrMessageEvent) -> None: """Queue a changed automatic UMO name for background persistence. diff --git a/tests/unit/test_event_bus.py b/tests/unit/test_event_bus.py index 8ccaa66135..8abb9605a6 100644 --- a/tests/unit/test_event_bus.py +++ b/tests/unit/test_event_bus.py @@ -113,8 +113,9 @@ async def test_dispatch_coalesces_auto_name_changes( processed = asyncio.Event() processed_count = 0 - async def execute_and_count(event): # noqa: ARG001 + async def execute_and_count(event): nonlocal processed_count + event.is_wake = True processed_count += 1 if processed_count == 3: processed.set() @@ -176,7 +177,8 @@ async def block_database_write(**kwargs): # noqa: ARG001 database_started.set() await release_database.wait() - async def execute_and_signal(event): # noqa: ARG001 + async def execute_and_signal(event): + event.is_wake = True pipeline_started.set() db_helper = MagicMock() @@ -214,6 +216,53 @@ async def execute_and_signal(event): # noqa: ARG001 with suppress(asyncio.CancelledError): await task + @pytest.mark.asyncio + async def test_dispatch_skips_auto_name_when_event_does_not_wake( + self, + event_queue, + mock_pipeline_scheduler, + mock_config_manager, + ): + """Do not persist names from events ignored by the waking stage.""" + processed = asyncio.Event() + + async def execute_without_waking(event): + event.is_wake = False + processed.set() + + mock_pipeline_scheduler.execute.side_effect = execute_without_waking + db_helper = MagicMock() + db_helper.upsert_umo_auto_name = AsyncMock() + bus = EventBus( + event_queue=event_queue, + pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, + astrbot_config_mgr=mock_config_manager, + db_helper=db_helper, + ) + mock_event = MagicMock() + mock_event.unified_msg_origin = "test-platform:GroupMessage:group-1" + mock_event.message_obj = SimpleNamespace( + group=SimpleNamespace(group_name="Engineering Group") + ) + mock_event.get_group_id.return_value = "group-1" + mock_event.get_platform_id.return_value = "test-platform" + mock_event.get_platform_name.return_value = "Test Platform" + mock_event.get_sender_name.return_value = "Alice" + mock_event.get_sender_id.return_value = "sender-1" + mock_event.get_message_outline.return_value = "Hello" + await event_queue.put(mock_event) + + task = asyncio.create_task(bus.dispatch()) + try: + await asyncio.wait_for(processed.wait(), timeout=1.0) + await asyncio.sleep(0) + finally: + task.cancel() + with suppress(asyncio.CancelledError): + await task + + db_helper.upsert_umo_auto_name.assert_not_awaited() + @pytest.mark.asyncio async def test_dispatch_bounds_auto_name_cache( self, @@ -225,8 +274,9 @@ async def test_dispatch_bounds_auto_name_cache( processed = asyncio.Event() processed_count = 0 - async def execute_and_count(event): # noqa: ARG001 + async def execute_and_count(event): nonlocal processed_count + event.is_wake = True processed_count += 1 if processed_count == 3: processed.set() @@ -329,8 +379,9 @@ async def test_dispatch_skips_missing_group_and_sender_names( processed = asyncio.Event() processed_count = 0 - async def execute_and_count(event): # noqa: ARG001 + async def execute_and_count(event): nonlocal processed_count + event.is_wake = True processed_count += 1 if processed_count == 2: processed.set() From 69e64bb28d0976ea01384deea0ba454c586472c7 Mon Sep 17 00:00:00 2001 From: Soulter <905617992@qq.com> Date: Tue, 1 Sep 2026 23:13:34 +0800 Subject: [PATCH 3/4] refactor: record UMO names during wake checks --- astrbot/core/core_lifecycle.py | 5 +- astrbot/core/event_bus.py | 103 +----- astrbot/core/pipeline/context.py | 2 + astrbot/core/pipeline/waking_check/stage.py | 103 +++++- tests/unit/test_event_bus.py | 330 ------------------ tests/unit/test_waking_check_api_key_admin.py | 5 +- tests/unit/test_waking_check_umo_alias.py | 223 ++++++++++++ 7 files changed, 332 insertions(+), 439 deletions(-) create mode 100644 tests/unit/test_waking_check_umo_alias.py diff --git a/astrbot/core/core_lifecycle.py b/astrbot/core/core_lifecycle.py index a632964792..9961a8c52b 100644 --- a/astrbot/core/core_lifecycle.py +++ b/astrbot/core/core_lifecycle.py @@ -278,7 +278,6 @@ async def initialize(self) -> None: self.event_queue, self.pipeline_scheduler_mapping, self.astrbot_config_mgr, - self.db, ) # 记录启动时间 @@ -463,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 @@ -480,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/event_bus.py b/astrbot/core/event_bus.py index 043c956833..aa7b80f937 100644 --- a/astrbot/core/event_bus.py +++ b/astrbot/core/event_bus.py @@ -12,18 +12,13 @@ import asyncio from asyncio import Queue -from collections import OrderedDict from astrbot.core import logger from astrbot.core.astrbot_config_mgr import AstrBotConfigManager -from astrbot.core.db import BaseDatabase from astrbot.core.pipeline.scheduler import PipelineScheduler -from astrbot.core.umo_alias import get_event_auto_name from .platform import AstrMessageEvent -MAX_UMO_AUTO_NAME_CACHE_SIZE = 10_000 - class EventBus: """用于处理事件的分发和处理""" @@ -33,16 +28,11 @@ def __init__( event_queue: Queue, pipeline_scheduler_mapping: dict[str, PipelineScheduler], astrbot_config_mgr: AstrBotConfigManager, - db_helper: BaseDatabase | None = None, ) -> None: self.event_queue = event_queue # 事件队列 # abconf uuid -> scheduler self.pipeline_scheduler_mapping = pipeline_scheduler_mapping self.astrbot_config_mgr = astrbot_config_mgr - self.db_helper = db_helper - self._umo_auto_name_cache: OrderedDict[str, str] = OrderedDict() - self._pending_umo_auto_names: OrderedDict[str, tuple[str, str]] = OrderedDict() - self._umo_auto_name_writer_task: asyncio.Task[None] | None = None # 持有正在执行的 pipeline 任务的强引用, 防止 task 在 pending 状态被 GC 回收 self._pending_tasks: set[asyncio.Task] = set() @@ -59,101 +49,10 @@ async def dispatch(self) -> None: f"PipelineScheduler not found for id: {conf_id}, event ignored." ) continue - task = asyncio.create_task(self._execute_pipeline(scheduler, event)) - self._pending_tasks.add(task) - task.add_done_callback(self._on_task_done) - - async def _execute_pipeline( - self, - scheduler: PipelineScheduler, - event: AstrMessageEvent, - ) -> None: - """Execute the pipeline and record the UMO name after a successful wake. - - Args: - scheduler: Pipeline scheduler selected for the event configuration. - event: Inbound platform event to process. - """ - await scheduler.execute(event) - if event.is_wake: - self._schedule_umo_auto_name_recording(event) - - def _schedule_umo_auto_name_recording(self, event: AstrMessageEvent) -> None: - """Queue a changed automatic UMO name for background persistence. - - Args: - event: Inbound platform 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._umo_auto_name_cache.get(umo) == auto_name: - self._umo_auto_name_cache.move_to_end(umo) - return - - self._umo_auto_name_cache[umo] = auto_name - self._umo_auto_name_cache.move_to_end(umo) - if len(self._umo_auto_name_cache) > MAX_UMO_AUTO_NAME_CACHE_SIZE: - self._umo_auto_name_cache.popitem(last=False) - - self._pending_umo_auto_names[umo] = ( - str(event.get_sender_id() or ""), - auto_name, - ) - self._pending_umo_auto_names.move_to_end(umo) - if len(self._pending_umo_auto_names) > MAX_UMO_AUTO_NAME_CACHE_SIZE: - dropped_umo, (_, dropped_name) = self._pending_umo_auto_names.popitem( - last=False - ) - if self._umo_auto_name_cache.get(dropped_umo) == dropped_name: - self._umo_auto_name_cache.pop(dropped_umo, None) - - if ( - self._umo_auto_name_writer_task is None - or self._umo_auto_name_writer_task.done() - ): - task = asyncio.create_task( - self._flush_umo_auto_names(), - name="umo_auto_name_writer", - ) - self._umo_auto_name_writer_task = task + task = asyncio.create_task(scheduler.execute(event)) self._pending_tasks.add(task) task.add_done_callback(self._on_task_done) - async def _flush_umo_auto_names(self) -> None: - """Persist queued UMO names sequentially, coalescing changes per UMO.""" - if self.db_helper is None: - return - - try: - while self._pending_umo_auto_names: - umo, (creator_sender_id, auto_name) = ( - self._pending_umo_auto_names.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_umo_auto_names - and self._umo_auto_name_cache.get(umo) == auto_name - ): - self._umo_auto_name_cache.pop(umo, None) - finally: - self._umo_auto_name_writer_task = None - def _on_task_done(self, task: asyncio.Task) -> None: """pipeline 任务结束回调: 移除强引用并暴露未捕获的异常""" self._pending_tasks.discard(task) 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..3d09cb8841 100644 --- a/astrbot/core/pipeline/waking_check/stage.py +++ b/astrbot/core/pipeline/waking_check/stage.py @@ -1,3 +1,5 @@ +import asyncio +from collections import OrderedDict from collections.abc import AsyncGenerator, Callable from astrbot import logger @@ -10,10 +12,13 @@ from astrbot.core.star.session_plugin_manager import SessionPluginManager from astrbot.core.star.star import star_map from astrbot.core.star.star_handler import EventType, star_handlers_registry +from astrbot.core.umo_alias import get_event_auto_name from ..context import PipelineContext from ..stage import Stage, register_stage +MAX_UMO_AUTO_NAME_CACHE_SIZE = 10_000 + UNIQUE_SESSION_ID_BUILDERS: dict[str, Callable[[AstrMessageEvent], str | None]] = { "aiocqhttp": lambda e: f"{e.get_sender_id()}_{e.get_group_id()}", "slack": lambda e: f"{e.get_sender_id()}_{e.get_group_id()}", @@ -73,6 +78,98 @@ 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.db_helper = ctx.db_helper + self._umo_auto_name_cache: OrderedDict[str, str] = OrderedDict() + self._pending_umo_auto_names: OrderedDict[str, tuple[str, str]] = OrderedDict() + self._umo_auto_name_writer_task: asyncio.Task[None] | None = None + + def _schedule_umo_auto_name_recording(self, event: AstrMessageEvent) -> None: + """Queue a changed automatic UMO name for background persistence. + + 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._umo_auto_name_cache.get(umo) == auto_name: + self._umo_auto_name_cache.move_to_end(umo) + return + + self._umo_auto_name_cache[umo] = auto_name + self._umo_auto_name_cache.move_to_end(umo) + if len(self._umo_auto_name_cache) > MAX_UMO_AUTO_NAME_CACHE_SIZE: + self._umo_auto_name_cache.popitem(last=False) + + self._pending_umo_auto_names[umo] = ( + str(event.get_sender_id() or ""), + auto_name, + ) + self._pending_umo_auto_names.move_to_end(umo) + if len(self._pending_umo_auto_names) > MAX_UMO_AUTO_NAME_CACHE_SIZE: + dropped_umo, (_, dropped_name) = self._pending_umo_auto_names.popitem( + last=False + ) + if self._umo_auto_name_cache.get(dropped_umo) == dropped_name: + self._umo_auto_name_cache.pop(dropped_umo, None) + + if ( + self._umo_auto_name_writer_task is None + or self._umo_auto_name_writer_task.done() + ): + task = asyncio.create_task( + self._flush_umo_auto_names(), + name=f"umo_auto_name_writer:{self.ctx.astrbot_config_id}", + ) + self._umo_auto_name_writer_task = task + task.add_done_callback(self._on_umo_auto_name_writer_done) + + async def _flush_umo_auto_names(self) -> None: + """Persist queued UMO names sequentially, coalescing changes per UMO.""" + if self.db_helper is None: + return + + try: + while self._pending_umo_auto_names: + umo, (creator_sender_id, auto_name) = ( + self._pending_umo_auto_names.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_umo_auto_names + and self._umo_auto_name_cache.get(umo) == auto_name + ): + self._umo_auto_name_cache.pop(umo, None) + finally: + self._umo_auto_name_writer_task = None + + @staticmethod + def _on_umo_auto_name_writer_done(task: asyncio.Task[None]) -> None: + """Expose unexpected automatic-name writer failures. + + Args: + task: Completed writer task. + """ + if task.cancelled(): + return + exc = task.exception() + if exc is not None: + logger.error("UMO automatic-name writer failed.", exc_info=exc) async def process( self, @@ -218,6 +315,8 @@ async def process( f"{star_map[handler.handler_module_path].name}.", ) event.stop_event() + if event.is_wake: + self._schedule_umo_auto_name_recording(event) return is_wake = True @@ -244,5 +343,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._schedule_umo_auto_name_recording(event) + else: event.stop_event() diff --git a/tests/unit/test_event_bus.py b/tests/unit/test_event_bus.py index 8abb9605a6..1ecdbf1e31 100644 --- a/tests/unit/test_event_bus.py +++ b/tests/unit/test_event_bus.py @@ -2,7 +2,6 @@ import asyncio from contextlib import suppress -from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -102,335 +101,6 @@ async def execute_and_signal(event): # noqa: ARG001 "test-platform:group:123" ) - @pytest.mark.asyncio - async def test_dispatch_coalesces_auto_name_changes( - self, - event_queue, - mock_pipeline_scheduler, - mock_config_manager, - ): - """Persist only the latest automatic name from a queued event burst.""" - processed = asyncio.Event() - processed_count = 0 - - async def execute_and_count(event): - nonlocal processed_count - event.is_wake = True - processed_count += 1 - if processed_count == 3: - processed.set() - - mock_pipeline_scheduler.execute.side_effect = execute_and_count - db_helper = MagicMock() - db_helper.upsert_umo_auto_name = AsyncMock() - bus = EventBus( - event_queue=event_queue, - pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, - astrbot_config_mgr=mock_config_manager, - db_helper=db_helper, - ) - - for group_name in ("Engineering Group", "Engineering Group", "Renamed"): - mock_event = MagicMock() - mock_event.unified_msg_origin = "test-platform:GroupMessage:group-1" - mock_event.message_obj = SimpleNamespace( - group=SimpleNamespace(group_name=group_name) - ) - mock_event.get_group_id.return_value = "group-1" - mock_event.get_platform_id.return_value = "test-platform" - mock_event.get_platform_name.return_value = "Test Platform" - mock_event.get_sender_name.return_value = "Alice" - mock_event.get_sender_id.return_value = "sender-1" - mock_event.get_message_outline.return_value = "Hello" - await event_queue.put(mock_event) - - task = asyncio.create_task(bus.dispatch()) - try: - await asyncio.wait_for(processed.wait(), timeout=1.0) - finally: - task.cancel() - with suppress(asyncio.CancelledError): - await task - - writer_task = bus._umo_auto_name_writer_task - if writer_task is not None: - await writer_task - assert db_helper.upsert_umo_auto_name.await_count == 1 - assert [ - call.kwargs["auto_name"] - for call in db_helper.upsert_umo_auto_name.await_args_list - ] == ["Renamed"] - - @pytest.mark.asyncio - async def test_dispatch_does_not_wait_for_auto_name_database_write( - self, - event_queue, - mock_pipeline_scheduler, - mock_config_manager, - ): - """Start the pipeline while the background alias writer is blocked.""" - database_started = asyncio.Event() - release_database = asyncio.Event() - pipeline_started = asyncio.Event() - - async def block_database_write(**kwargs): # noqa: ARG001 - database_started.set() - await release_database.wait() - - async def execute_and_signal(event): - event.is_wake = True - pipeline_started.set() - - db_helper = MagicMock() - db_helper.upsert_umo_auto_name = AsyncMock(side_effect=block_database_write) - mock_pipeline_scheduler.execute.side_effect = execute_and_signal - bus = EventBus( - event_queue=event_queue, - pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, - astrbot_config_mgr=mock_config_manager, - db_helper=db_helper, - ) - mock_event = MagicMock() - mock_event.unified_msg_origin = "test-platform:GroupMessage:group-1" - mock_event.message_obj = SimpleNamespace( - group=SimpleNamespace(group_name="Engineering Group") - ) - mock_event.get_group_id.return_value = "group-1" - mock_event.get_platform_id.return_value = "test-platform" - mock_event.get_platform_name.return_value = "Test Platform" - mock_event.get_sender_name.return_value = "Alice" - mock_event.get_sender_id.return_value = "sender-1" - mock_event.get_message_outline.return_value = "Hello" - await event_queue.put(mock_event) - - task = asyncio.create_task(bus.dispatch()) - try: - await asyncio.wait_for(database_started.wait(), timeout=1.0) - await asyncio.wait_for(pipeline_started.wait(), timeout=1.0) - finally: - release_database.set() - writer_task = bus._umo_auto_name_writer_task - if writer_task is not None: - await writer_task - task.cancel() - with suppress(asyncio.CancelledError): - await task - - @pytest.mark.asyncio - async def test_dispatch_skips_auto_name_when_event_does_not_wake( - self, - event_queue, - mock_pipeline_scheduler, - mock_config_manager, - ): - """Do not persist names from events ignored by the waking stage.""" - processed = asyncio.Event() - - async def execute_without_waking(event): - event.is_wake = False - processed.set() - - mock_pipeline_scheduler.execute.side_effect = execute_without_waking - db_helper = MagicMock() - db_helper.upsert_umo_auto_name = AsyncMock() - bus = EventBus( - event_queue=event_queue, - pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, - astrbot_config_mgr=mock_config_manager, - db_helper=db_helper, - ) - mock_event = MagicMock() - mock_event.unified_msg_origin = "test-platform:GroupMessage:group-1" - mock_event.message_obj = SimpleNamespace( - group=SimpleNamespace(group_name="Engineering Group") - ) - mock_event.get_group_id.return_value = "group-1" - mock_event.get_platform_id.return_value = "test-platform" - mock_event.get_platform_name.return_value = "Test Platform" - mock_event.get_sender_name.return_value = "Alice" - mock_event.get_sender_id.return_value = "sender-1" - mock_event.get_message_outline.return_value = "Hello" - await event_queue.put(mock_event) - - task = asyncio.create_task(bus.dispatch()) - try: - await asyncio.wait_for(processed.wait(), timeout=1.0) - await asyncio.sleep(0) - finally: - task.cancel() - with suppress(asyncio.CancelledError): - await task - - db_helper.upsert_umo_auto_name.assert_not_awaited() - - @pytest.mark.asyncio - async def test_dispatch_bounds_auto_name_cache( - self, - event_queue, - mock_pipeline_scheduler, - mock_config_manager, - ): - """Evict the least recently used UMO when the cache reaches its bound.""" - processed = asyncio.Event() - processed_count = 0 - - async def execute_and_count(event): - nonlocal processed_count - event.is_wake = True - processed_count += 1 - if processed_count == 3: - processed.set() - - mock_pipeline_scheduler.execute.side_effect = execute_and_count - db_helper = MagicMock() - db_helper.upsert_umo_auto_name = AsyncMock() - with patch("astrbot.core.event_bus.MAX_UMO_AUTO_NAME_CACHE_SIZE", 2): - bus = EventBus( - event_queue=event_queue, - pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, - astrbot_config_mgr=mock_config_manager, - db_helper=db_helper, - ) - - for index in range(3): - mock_event = MagicMock() - mock_event.unified_msg_origin = ( - f"test-platform:GroupMessage:group-{index}" - ) - mock_event.message_obj = SimpleNamespace( - group=SimpleNamespace(group_name=f"Group {index}") - ) - mock_event.get_group_id.return_value = f"group-{index}" - mock_event.get_platform_id.return_value = "test-platform" - mock_event.get_platform_name.return_value = "Test Platform" - mock_event.get_sender_name.return_value = "Alice" - mock_event.get_sender_id.return_value = "sender-1" - mock_event.get_message_outline.return_value = "Hello" - await event_queue.put(mock_event) - - task = asyncio.create_task(bus.dispatch()) - try: - await asyncio.wait_for(processed.wait(), timeout=1.0) - finally: - task.cancel() - with suppress(asyncio.CancelledError): - await task - - assert list(bus._umo_auto_name_cache) == [ - "test-platform:GroupMessage:group-1", - "test-platform:GroupMessage:group-2", - ] - assert not bus._pending_umo_auto_names - assert db_helper.upsert_umo_auto_name.await_count == 2 - - @pytest.mark.asyncio - async def test_auto_name_writer_retries_after_database_failure( - self, - event_queue, - mock_pipeline_scheduler, - mock_config_manager, - ): - """Evict a failed optimistic cache entry so the next event retries.""" - db_helper = MagicMock() - db_helper.upsert_umo_auto_name = AsyncMock( - side_effect=[RuntimeError("database unavailable"), None] - ) - bus = EventBus( - event_queue=event_queue, - pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, - astrbot_config_mgr=mock_config_manager, - db_helper=db_helper, - ) - mock_event = MagicMock() - mock_event.unified_msg_origin = "test-platform:GroupMessage:group-1" - mock_event.message_obj = SimpleNamespace( - group=SimpleNamespace(group_name="Engineering Group") - ) - mock_event.get_group_id.return_value = "group-1" - mock_event.get_sender_id.return_value = "sender-1" - - with patch("astrbot.core.event_bus.logger"): - bus._schedule_umo_auto_name_recording(mock_event) - first_writer = bus._umo_auto_name_writer_task - assert first_writer is not None - await first_writer - - assert mock_event.unified_msg_origin not in bus._umo_auto_name_cache - - bus._schedule_umo_auto_name_recording(mock_event) - second_writer = bus._umo_auto_name_writer_task - assert second_writer is not None - await second_writer - - assert db_helper.upsert_umo_auto_name.await_count == 2 - assert ( - bus._umo_auto_name_cache[mock_event.unified_msg_origin] - == "Engineering Group" - ) - - @pytest.mark.asyncio - async def test_dispatch_skips_missing_group_and_sender_names( - self, - event_queue, - mock_pipeline_scheduler, - mock_config_manager, - ): - """Do not persist ID fallbacks as automatic UMO names.""" - processed = asyncio.Event() - processed_count = 0 - - async def execute_and_count(event): - nonlocal processed_count - event.is_wake = True - processed_count += 1 - if processed_count == 2: - processed.set() - - mock_pipeline_scheduler.execute.side_effect = execute_and_count - db_helper = MagicMock() - db_helper.upsert_umo_auto_name = AsyncMock() - bus = EventBus( - event_queue=event_queue, - pipeline_scheduler_mapping={"test-conf-id": mock_pipeline_scheduler}, - astrbot_config_mgr=mock_config_manager, - db_helper=db_helper, - ) - - group_event = MagicMock() - group_event.unified_msg_origin = "test-platform:GroupMessage:group-1" - group_event.message_obj = SimpleNamespace( - group=SimpleNamespace(group_name=None) - ) - group_event.get_group_id.return_value = "group-1" - group_event.get_sender_name.return_value = "Alice" - group_event.get_sender_id.return_value = "sender-1" - group_event.get_platform_id.return_value = "test-platform" - group_event.get_platform_name.return_value = "Test Platform" - group_event.get_message_outline.return_value = "Hello" - - 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" - friend_event.get_platform_id.return_value = "test-platform" - friend_event.get_platform_name.return_value = "Test Platform" - friend_event.get_message_outline.return_value = "Hello" - - await event_queue.put(group_event) - await event_queue.put(friend_event) - task = asyncio.create_task(bus.dispatch()) - try: - await asyncio.wait_for(processed.wait(), timeout=1.0) - finally: - task.cancel() - with suppress(asyncio.CancelledError): - await task - - db_helper.upsert_umo_auto_name.assert_not_awaited() - assert not bus._umo_auto_name_cache - @pytest.mark.asyncio async def test_dispatch_handles_missing_scheduler( self, diff --git a/tests/unit/test_waking_check_api_key_admin.py b/tests/unit/test_waking_check_api_key_admin.py index 0d85ce22fe..25ea608f1c 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.db_helper = None 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..aacb84e8df --- /dev/null +++ b/tests/unit/test_waking_check_umo_alias.py @@ -0,0 +1,223 @@ +"""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.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_writer_task is None + + awakened_event = make_group_event("group-1", "Engineering") + await stage.process(awakened_event) + writer_task = stage._umo_auto_name_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() + stage = await make_stage(db_helper) + + for group_name in ("Engineering", "Engineering", "Renamed"): + stage._schedule_umo_auto_name_recording(make_group_event("group-1", group_name)) + + writer_task = stage._umo_auto_name_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() + stage = await make_stage(db_helper) + + stage._schedule_umo_auto_name_recording(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" + stage._schedule_umo_auto_name_recording(friend_event) + + assert stage._umo_auto_name_writer_task is None + assert not stage._umo_auto_name_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() + stage = await make_stage(db_helper) + + with patch( + "astrbot.core.pipeline.waking_check.stage.MAX_UMO_AUTO_NAME_CACHE_SIZE", + 2, + ): + for index in range(3): + stage._schedule_umo_auto_name_recording( + make_group_event(f"group-{index}", f"Group {index}") + ) + + writer_task = stage._umo_auto_name_writer_task + assert writer_task is not None + await writer_task + + assert list(stage._umo_auto_name_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] + ) + stage = await make_stage(db_helper) + event = make_group_event("group-1", "Engineering") + + with patch("astrbot.core.pipeline.waking_check.stage.logger"): + stage._schedule_umo_auto_name_recording(event) + first_writer = stage._umo_auto_name_writer_task + assert first_writer is not None + await first_writer + + assert event.unified_msg_origin not in stage._umo_auto_name_cache + + stage._schedule_umo_auto_name_recording(event) + second_writer = stage._umo_auto_name_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) + stage = await make_stage(db_helper) + stage._schedule_umo_auto_name_recording(make_group_event("group-1", "Engineering")) + + await asyncio.wait_for(database_started.wait(), timeout=1.0) + release_database.set() + writer_task = stage._umo_auto_name_writer_task + if writer_task is not None: + await writer_task From bfc178948d2a9f08efe682df5631ea6ea0748374 Mon Sep 17 00:00:00 2001 From: Soulter <905617992@qq.com> Date: Tue, 1 Sep 2026 23:19:58 +0800 Subject: [PATCH 4/4] refactor: extract UMO auto-name recorder --- astrbot/core/pipeline/waking_check/stage.py | 104 +---------------- .../pipeline/waking_check/umo_auto_name.py | 110 ++++++++++++++++++ tests/unit/test_waking_check_api_key_admin.py | 2 +- tests/unit/test_waking_check_umo_alias.py | 53 +++++---- 4 files changed, 143 insertions(+), 126 deletions(-) create mode 100644 astrbot/core/pipeline/waking_check/umo_auto_name.py diff --git a/astrbot/core/pipeline/waking_check/stage.py b/astrbot/core/pipeline/waking_check/stage.py index 3d09cb8841..f02915bd64 100644 --- a/astrbot/core/pipeline/waking_check/stage.py +++ b/astrbot/core/pipeline/waking_check/stage.py @@ -1,5 +1,3 @@ -import asyncio -from collections import OrderedDict from collections.abc import AsyncGenerator, Callable from astrbot import logger @@ -12,12 +10,10 @@ from astrbot.core.star.session_plugin_manager import SessionPluginManager from astrbot.core.star.star import star_map from astrbot.core.star.star_handler import EventType, star_handlers_registry -from astrbot.core.umo_alias import get_event_auto_name from ..context import PipelineContext from ..stage import Stage, register_stage - -MAX_UMO_AUTO_NAME_CACHE_SIZE = 10_000 +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()}", @@ -78,98 +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.db_helper = ctx.db_helper - self._umo_auto_name_cache: OrderedDict[str, str] = OrderedDict() - self._pending_umo_auto_names: OrderedDict[str, tuple[str, str]] = OrderedDict() - self._umo_auto_name_writer_task: asyncio.Task[None] | None = None - - def _schedule_umo_auto_name_recording(self, event: AstrMessageEvent) -> None: - """Queue a changed automatic UMO name for background persistence. - - 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._umo_auto_name_cache.get(umo) == auto_name: - self._umo_auto_name_cache.move_to_end(umo) - return - - self._umo_auto_name_cache[umo] = auto_name - self._umo_auto_name_cache.move_to_end(umo) - if len(self._umo_auto_name_cache) > MAX_UMO_AUTO_NAME_CACHE_SIZE: - self._umo_auto_name_cache.popitem(last=False) - - self._pending_umo_auto_names[umo] = ( - str(event.get_sender_id() or ""), - auto_name, + self._umo_auto_name_recorder = UmoAutoNameRecorder( + ctx.db_helper, + ctx.astrbot_config_id, ) - self._pending_umo_auto_names.move_to_end(umo) - if len(self._pending_umo_auto_names) > MAX_UMO_AUTO_NAME_CACHE_SIZE: - dropped_umo, (_, dropped_name) = self._pending_umo_auto_names.popitem( - last=False - ) - if self._umo_auto_name_cache.get(dropped_umo) == dropped_name: - self._umo_auto_name_cache.pop(dropped_umo, None) - - if ( - self._umo_auto_name_writer_task is None - or self._umo_auto_name_writer_task.done() - ): - task = asyncio.create_task( - self._flush_umo_auto_names(), - name=f"umo_auto_name_writer:{self.ctx.astrbot_config_id}", - ) - self._umo_auto_name_writer_task = task - task.add_done_callback(self._on_umo_auto_name_writer_done) - - async def _flush_umo_auto_names(self) -> None: - """Persist queued UMO names sequentially, coalescing changes per UMO.""" - if self.db_helper is None: - return - - try: - while self._pending_umo_auto_names: - umo, (creator_sender_id, auto_name) = ( - self._pending_umo_auto_names.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_umo_auto_names - and self._umo_auto_name_cache.get(umo) == auto_name - ): - self._umo_auto_name_cache.pop(umo, None) - finally: - self._umo_auto_name_writer_task = None - - @staticmethod - def _on_umo_auto_name_writer_done(task: asyncio.Task[None]) -> None: - """Expose unexpected automatic-name writer failures. - - Args: - task: Completed writer task. - """ - if task.cancelled(): - return - exc = task.exception() - if exc is not None: - logger.error("UMO automatic-name writer failed.", exc_info=exc) async def process( self, @@ -316,7 +224,7 @@ async def process( ) event.stop_event() if event.is_wake: - self._schedule_umo_auto_name_recording(event) + self._umo_auto_name_recorder.schedule(event) return is_wake = True @@ -344,6 +252,6 @@ async def process( event.set_extra("handlers_parsed_params", handlers_parsed_params) if is_wake: - self._schedule_umo_auto_name_recording(event) + 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/tests/unit/test_waking_check_api_key_admin.py b/tests/unit/test_waking_check_api_key_admin.py index 25ea608f1c..ee009c1e97 100644 --- a/tests/unit/test_waking_check_api_key_admin.py +++ b/tests/unit/test_waking_check_api_key_admin.py @@ -39,7 +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.db_helper = None + stage._umo_auto_name_recorder = MagicMock() event = MagicMock() event.message_str = "hello" diff --git a/tests/unit/test_waking_check_umo_alias.py b/tests/unit/test_waking_check_umo_alias.py index aacb84e8df..b74f9cc2cf 100644 --- a/tests/unit/test_waking_check_umo_alias.py +++ b/tests/unit/test_waking_check_umo_alias.py @@ -7,6 +7,7 @@ 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 @@ -90,11 +91,11 @@ async def return_handlers(_event, handlers): ignored_event = make_group_event("group-1", "Engineering", "hello") await stage.process(ignored_event) - assert stage._umo_auto_name_writer_task is None + 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_writer_task + writer_task = stage._umo_auto_name_recorder._writer_task assert writer_task is not None await writer_task @@ -110,12 +111,12 @@ 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() - stage = await make_stage(db_helper) + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") for group_name in ("Engineering", "Engineering", "Renamed"): - stage._schedule_umo_auto_name_recording(make_group_event("group-1", group_name)) + recorder.schedule(make_group_event("group-1", group_name)) - writer_task = stage._umo_auto_name_writer_task + writer_task = recorder._writer_task assert writer_task is not None await writer_task @@ -131,9 +132,9 @@ 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() - stage = await make_stage(db_helper) + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") - stage._schedule_umo_auto_name_recording(make_group_event("group-1", None)) + recorder.schedule(make_group_event("group-1", None)) friend_event = MagicMock() friend_event.unified_msg_origin = "test-platform:FriendMessage:sender-2" @@ -141,10 +142,10 @@ async def test_waking_stage_skips_missing_group_and_sender_names(): friend_event.get_group_id.return_value = "" friend_event.get_sender_name.return_value = "" friend_event.get_sender_id.return_value = "sender-2" - stage._schedule_umo_auto_name_recording(friend_event) + recorder.schedule(friend_event) - assert stage._umo_auto_name_writer_task is None - assert not stage._umo_auto_name_cache + assert recorder._writer_task is None + assert not recorder._cache db_helper.upsert_umo_auto_name.assert_not_awaited() @@ -153,22 +154,20 @@ 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() - stage = await make_stage(db_helper) + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") with patch( - "astrbot.core.pipeline.waking_check.stage.MAX_UMO_AUTO_NAME_CACHE_SIZE", + "astrbot.core.pipeline.waking_check.umo_auto_name.MAX_UMO_AUTO_NAME_CACHE_SIZE", 2, ): for index in range(3): - stage._schedule_umo_auto_name_recording( - make_group_event(f"group-{index}", f"Group {index}") - ) + recorder.schedule(make_group_event(f"group-{index}", f"Group {index}")) - writer_task = stage._umo_auto_name_writer_task + writer_task = recorder._writer_task assert writer_task is not None await writer_task - assert list(stage._umo_auto_name_cache) == [ + assert list(recorder._cache) == [ "test-platform:GroupMessage:group-1", "test-platform:GroupMessage:group-2", ] @@ -182,19 +181,19 @@ async def test_waking_stage_retries_after_database_failure(): db_helper.upsert_umo_auto_name = AsyncMock( side_effect=[RuntimeError("database unavailable"), None] ) - stage = await make_stage(db_helper) + recorder = UmoAutoNameRecorder(db_helper, "test-conf-id") event = make_group_event("group-1", "Engineering") - with patch("astrbot.core.pipeline.waking_check.stage.logger"): - stage._schedule_umo_auto_name_recording(event) - first_writer = stage._umo_auto_name_writer_task + 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 stage._umo_auto_name_cache + assert event.unified_msg_origin not in recorder._cache - stage._schedule_umo_auto_name_recording(event) - second_writer = stage._umo_auto_name_writer_task + recorder.schedule(event) + second_writer = recorder._writer_task assert second_writer is not None await second_writer @@ -213,11 +212,11 @@ async def block_database_write(**kwargs): # noqa: ARG001 db_helper = MagicMock() db_helper.upsert_umo_auto_name = AsyncMock(side_effect=block_database_write) - stage = await make_stage(db_helper) - stage._schedule_umo_auto_name_recording(make_group_event("group-1", "Engineering")) + 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 = stage._umo_auto_name_writer_task + writer_task = recorder._writer_task if writer_task is not None: await writer_task