diff --git a/astrbot/core/auth/admission.py b/astrbot/core/auth/admission.py new file mode 100644 index 0000000000..17dc1ab668 --- /dev/null +++ b/astrbot/core/auth/admission.py @@ -0,0 +1,319 @@ +"""Canonical admission identity keys and overlay composition. + +Session keys ignore unique-session rewrites of ``session_id``. Sender keys reuse +``Subject.im``. Composition is pure; persistence owners still choose the +preference scope they read. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import StrEnum +from typing import Protocol + +from astrbot.core.auth.models import Subject, normalize_subject_component +from astrbot.core.platform.message_type import MessageType + +SESSION_SERVICE_CONFIG_KEY = "session_service_config" + + +class UnlistedPolicy(StrEnum): + """Default for identities that have no explicit overlay.""" + + ALLOW = "allow" + DENY = "deny" + + +class ConversationKind(StrEnum): + """Admission conversation axis. Group vs private, not unique-session.""" + + GROUP = "group" + PRIVATE = "private" + + +class AdmissionEvent(Protocol): + """Inbound facts needed to mint admission keys without importing events.""" + + def get_platform_name(self) -> str: + raise NotImplementedError + + def get_sender_id(self) -> str: + raise NotImplementedError + + def get_self_id(self) -> str: + raise NotImplementedError + + def get_group_id(self) -> str: + raise NotImplementedError + + def get_session_id(self) -> str: + raise NotImplementedError + + def get_message_type(self) -> MessageType | str: + raise NotImplementedError + + +@dataclass(frozen=True, slots=True) +class SessionAdmissionOverlay: + """UMO-scoped admission fields. ``None`` means unwritten.""" + + session_enabled: bool | None = None + session_blocked: bool = False + llm_enabled: bool | None = None + listed: bool = False + + +@dataclass(frozen=True, slots=True) +class SenderAdmissionOverlay: + """UID-scoped admission fields. ``None`` means follow the session.""" + + blocked: bool = False + llm_enabled: bool | None = None + listed: bool = False + + +@dataclass(frozen=True, slots=True) +class AdmissionDecision: + """Event and built-in LLM admission after overlay composition.""" + + admit_event: bool + admit_llm: bool + session_enabled: bool + session_blocked: bool + sender_blocked: bool + + +def platform_instance_from_event(event: AdmissionEvent) -> str: + """Return the platform instance id used in identity keys. + + Args: + event: Inbound event exposing platform accessors. + + Returns: + ``get_platform_id()`` when it is a non-empty string, otherwise the + platform name, otherwise ``unknown``. + """ + + get_platform_id = getattr(event, "get_platform_id", None) + platform_id = get_platform_id() if callable(get_platform_id) else None + if isinstance(platform_id, str) and platform_id.strip(): + return platform_id + name = event.get_platform_name() + if isinstance(name, str) and name.strip(): + return name + return "unknown" + + +def conversation_kind_from_event(event: AdmissionEvent) -> ConversationKind: + """Return group vs private from the inbound message type. + + Args: + event: Inbound event. + + Returns: + ``group`` only for ``MessageType.GROUP_MESSAGE``; every other type is + ``private``. + """ + + message_type = event.get_message_type() + if message_type is MessageType.GROUP_MESSAGE: + return ConversationKind.GROUP + if message_type == MessageType.GROUP_MESSAGE.value: + return ConversationKind.GROUP + return ConversationKind.PRIVATE + + +def session_admission_key( + *, + platform_instance: str, + conversation_kind: ConversationKind | str, + conversation_id: str, +) -> str: + """Build ``session:{platform_instance}:{group|private}:{id}``. + + Args: + platform_instance: Adapter instance id, not a display name when both + exist. + conversation_kind: ``group`` or ``private``. + conversation_id: Group id or private peer id. + + Returns: + Canonical session admission key with normalized components. + """ + + kind = ConversationKind(conversation_kind) + return ( + "session:" + f"{normalize_subject_component(platform_instance, 'platform instance')}:" + f"{kind.value}:" + f"{normalize_subject_component(conversation_id, 'conversation id')}" + ) + + +def session_admission_key_from_event(event: AdmissionEvent) -> str: + """Mint a session key that ignores unique-session ``session_id`` rewrites. + + Group messages always use ``get_group_id()``. Private messages use the + peer ``sender_id`` (falling back to ``session_id``). + + Args: + event: Inbound event. ``session_id`` may already be unique-session + rewritten. + + Returns: + Canonical session admission key. + """ + + kind = conversation_kind_from_event(event) + if kind is ConversationKind.GROUP: + conversation_id = str(event.get_group_id() or "").strip() or "unknown" + else: + conversation_id = ( + str(event.get_sender_id() or event.get_session_id() or "").strip() + or "unknown" + ) + return session_admission_key( + platform_instance=platform_instance_from_event(event), + conversation_kind=kind, + conversation_id=conversation_id, + ) + + +def sender_admission_key_from_event(event: AdmissionEvent) -> str: + """Return the sender key, reusing ``Subject.im`` when possible. + + Args: + event: Inbound event. An attached IM ``subject`` is preferred. + + Returns: + ``im:{platform_instance}:{bot_account_id}:{sender_id}``. + """ + + subject = getattr(event, "subject", None) + if isinstance(subject, Subject) and subject.kind == "im": + return subject.id + return Subject.im( + platform_instance=platform_instance_from_event(event), + bot_account_id=str(event.get_self_id() or "").strip() or "default", + sender_id=str(event.get_sender_id() or "").strip() or "unknown", + ).id + + +def session_overlay_from_config(config: object) -> SessionAdmissionOverlay: + """Parse a ``session_service_config`` mapping into a session overlay. + + Args: + config: Preference value, typically a dict. + + Returns: + Overlay with unwritten fields left as defaults. + """ + + if not isinstance(config, dict): + return SessionAdmissionOverlay() + session_enabled = config.get("session_enabled") + session_blocked = config.get("session_blocked") + llm_enabled = config.get("llm_enabled") + listed = any( + isinstance(config.get(key), bool) + for key in ("session_enabled", "session_blocked", "llm_enabled") + ) + return SessionAdmissionOverlay( + session_enabled=session_enabled if isinstance(session_enabled, bool) else None, + session_blocked=session_blocked if isinstance(session_blocked, bool) else False, + llm_enabled=llm_enabled if isinstance(llm_enabled, bool) else None, + listed=listed, + ) + + +def sender_overlay_from_config(config: object) -> SenderAdmissionOverlay: + """Parse a sender-scoped ``session_service_config`` mapping. + + Args: + config: Preference value, typically a dict. + + Returns: + Overlay with unwritten fields left as defaults. + """ + + if not isinstance(config, dict): + return SenderAdmissionOverlay() + blocked = config.get("blocked") + llm_enabled = config.get("llm_enabled") + listed = any( + isinstance(config.get(key), bool) for key in ("blocked", "llm_enabled") + ) + return SenderAdmissionOverlay( + blocked=blocked if isinstance(blocked, bool) else False, + llm_enabled=llm_enabled if isinstance(llm_enabled, bool) else None, + listed=listed, + ) + + +def composed_llm_enabled( + session: SessionAdmissionOverlay, + sender: SenderAdmissionOverlay, +) -> bool: + """Return the LLM overlay after sender specificity, ignoring event drops. + + Unwritten session LLM defaults to enabled. An unwritten sender follows the + session. A written sender value wins, including VIP enable over a disabled + session. + + Args: + session: UMO overlay. + sender: UID overlay. + + Returns: + Whether built-in LLM is enabled by overlays alone. + """ + + session_llm = True if session.llm_enabled is None else session.llm_enabled + if sender.llm_enabled is None: + return session_llm + return sender.llm_enabled + + +def compose_admission( + session: SessionAdmissionOverlay, + sender: SenderAdmissionOverlay, + *, + unlisted_sessions: UnlistedPolicy | str = UnlistedPolicy.ALLOW, + unlisted_senders: UnlistedPolicy | str = UnlistedPolicy.ALLOW, +) -> AdmissionDecision: + """Compose session and sender overlays into one admission decision. + + Refusal wins over allow. Sender LLM overlays are more specific than + session LLM overlays. ``session_blocked`` and sender ``blocked`` drop the + event and therefore the built-in LLM; a personal LLM exception cannot + revive a blocked session. + + Args: + session: UMO overlay. + sender: UID overlay. + unlisted_sessions: Policy when the session has no overlay. + unlisted_senders: Policy when the sender has no allow overlay. + + Returns: + Event and LLM admission plus the resolved session/sender flags. + """ + + session_policy = UnlistedPolicy(unlisted_sessions) + sender_policy = UnlistedPolicy(unlisted_senders) + session_enabled = ( + True if session.session_enabled is None else session.session_enabled + ) + session_blocked = session.session_blocked + sender_blocked = sender.blocked + session_denied = session_policy is UnlistedPolicy.DENY and not session.listed + sender_denied = sender_policy is UnlistedPolicy.DENY and not ( + sender.listed and not sender.blocked + ) + refused = session_blocked or sender_blocked or session_denied or sender_denied + return AdmissionDecision( + admit_event=not refused, + admit_llm=not refused and composed_llm_enabled(session, sender), + session_enabled=session_enabled, + session_blocked=session_blocked, + sender_blocked=sender_blocked, + ) diff --git a/astrbot/core/star/session_llm_manager.py b/astrbot/core/star/session_llm_manager.py index e500087820..8e3fa19eda 100644 --- a/astrbot/core/star/session_llm_manager.py +++ b/astrbot/core/star/session_llm_manager.py @@ -1,6 +1,13 @@ """会话服务管理器 - 负责管理每个会话的LLM、TTS等服务的启停状态""" from astrbot import logger +from astrbot.core.auth.admission import ( + SESSION_SERVICE_CONFIG_KEY, + composed_llm_enabled, + sender_admission_key_from_event, + sender_overlay_from_config, + session_overlay_from_config, +) from astrbot.core.platform.astr_message_event import AstrMessageEvent from astrbot.core.utils.shared_preferences import SharedPreferences @@ -11,6 +18,15 @@ class SessionServiceManager: def __init__(self, preferences: SharedPreferences) -> None: self.preferences = preferences + async def _service_config(self, scope: str, scope_id: str) -> dict: + config = await self.preferences.get_async( + scope=scope, + scope_id=scope_id, + key=SESSION_SERVICE_CONFIG_KEY, + default={}, + ) + return config if isinstance(config, dict) else {} + async def is_llm_enabled_for_session(self, session_id: str) -> bool: """检查LLM是否在指定会话中启用 @@ -21,21 +37,10 @@ async def is_llm_enabled_for_session(self, session_id: str) -> bool: bool: True表示启用,False表示禁用 """ - # 获取会话服务配置 - session_services = await self.preferences.get_async( - scope="umo", - scope_id=session_id, - key="session_service_config", - default={}, + overlay = session_overlay_from_config( + await self._service_config("umo", session_id) ) - - # 如果配置了该会话的LLM状态,返回该状态 - llm_enabled = session_services.get("llm_enabled") - if llm_enabled is not None: - return llm_enabled - - # 如果没有配置,默认为启用(兼容性考虑) - return True + return True if overlay.llm_enabled is None else overlay.llm_enabled async def set_llm_status_for_session(self, session_id: str, enabled: bool) -> None: """设置LLM在指定会话中的启停状态 @@ -45,26 +50,21 @@ async def set_llm_status_for_session(self, session_id: str, enabled: bool) -> No enabled: True表示启用,False表示禁用 """ - session_config = ( - await self.preferences.get_async( - scope="umo", - scope_id=session_id, - key="session_service_config", - default={}, - ) - or {} - ) + session_config = await self._service_config("umo", session_id) session_config["llm_enabled"] = enabled await self.preferences.put_async( scope="umo", scope_id=session_id, - key="session_service_config", + key=SESSION_SERVICE_CONFIG_KEY, value=session_config, ) async def should_process_llm_request(self, event: AstrMessageEvent) -> bool: """检查是否应该处理LLM请求 + Empty sender overlays follow the current UMO ``llm_enabled`` switch. + A written sender overlay is more specific, including VIP enable. + Args: event: 消息事件 @@ -72,8 +72,13 @@ async def should_process_llm_request(self, event: AstrMessageEvent) -> bool: bool: True表示应该处理,False表示跳过 """ - session_id = event.unified_msg_origin - return await self.is_llm_enabled_for_session(session_id) + session_overlay = session_overlay_from_config( + await self._service_config("umo", event.unified_msg_origin) + ) + sender_overlay = sender_overlay_from_config( + await self._service_config("sender", sender_admission_key_from_event(event)) + ) + return composed_llm_enabled(session_overlay, sender_overlay) # ============================================================================= # TTS 相关方法 @@ -89,21 +94,8 @@ async def is_tts_enabled_for_session(self, session_id: str) -> bool: bool: True表示启用,False表示禁用 """ - # 获取会话服务配置 - session_services = await self.preferences.get_async( - scope="umo", - scope_id=session_id, - key="session_service_config", - default={}, - ) - - # 如果配置了该会话的TTS状态,返回该状态 - tts_enabled = session_services.get("tts_enabled") - if tts_enabled is not None: - return tts_enabled - - # 如果没有配置,默认为启用(兼容性考虑) - return True + tts_enabled = (await self._service_config("umo", session_id)).get("tts_enabled") + return tts_enabled if isinstance(tts_enabled, bool) else True async def set_tts_status_for_session(self, session_id: str, enabled: bool) -> None: """设置TTS在指定会话中的启停状态 @@ -113,20 +105,12 @@ async def set_tts_status_for_session(self, session_id: str, enabled: bool) -> No enabled: True表示启用,False表示禁用 """ - session_config = ( - await self.preferences.get_async( - scope="umo", - scope_id=session_id, - key="session_service_config", - default={}, - ) - or {} - ) + session_config = await self._service_config("umo", session_id) session_config["tts_enabled"] = enabled await self.preferences.put_async( scope="umo", scope_id=session_id, - key="session_service_config", + key=SESSION_SERVICE_CONFIG_KEY, value=session_config, ) @@ -161,48 +145,41 @@ async def is_session_enabled(self, session_id: str) -> bool: bool: True表示启用,False表示禁用 """ - # 获取会话服务配置 - session_services = await self.preferences.get_async( - scope="umo", - scope_id=session_id, - key="session_service_config", - default={}, + overlay = session_overlay_from_config( + await self._service_config("umo", session_id) ) - - # 如果配置了该会话的整体状态,返回该状态 - session_enabled = session_services.get("session_enabled") - if session_enabled is not None: - return session_enabled - - # 如果没有配置,默认为启用(兼容性考虑) - return True + return True if overlay.session_enabled is None else overlay.session_enabled async def is_session_blocked(self, session_id: str) -> bool: """Check whether all functionality is blocked for a session.""" - session_services = await self.preferences.get_async( - scope="umo", - scope_id=session_id, - key="session_service_config", - default={}, + overlay = session_overlay_from_config( + await self._service_config("umo", session_id) ) - blocked = session_services.get("session_blocked") - return blocked if isinstance(blocked, bool) else False + return overlay.session_blocked + + async def is_sender_blocked(self, event: AstrMessageEvent) -> bool: + """Check whether the inbound sender has a UID ``blocked`` overlay. + + ``scope=sender`` is readable in this slice; missing rows are unblocked. + + Args: + event: Inbound event used to mint the sender key. + + Returns: + True when the sender overlay sets ``blocked``. + """ + overlay = sender_overlay_from_config( + await self._service_config("sender", sender_admission_key_from_event(event)) + ) + return overlay.blocked async def set_session_blocked(self, session_id: str, blocked: bool) -> None: """Block or unblock all functionality for a session.""" - session_config = ( - await self.preferences.get_async( - scope="umo", - scope_id=session_id, - key="session_service_config", - default={}, - ) - or {} - ) + session_config = await self._service_config("umo", session_id) session_config["session_blocked"] = blocked await self.preferences.put_async( scope="umo", scope_id=session_id, - key="session_service_config", + key=SESSION_SERVICE_CONFIG_KEY, value=session_config, ) diff --git a/tests/unit/platform/test_telegram_adapter.py b/tests/unit/platform/test_telegram_adapter.py index b9b3ccbda1..3ff54ddab4 100644 --- a/tests/unit/platform/test_telegram_adapter.py +++ b/tests/unit/platform/test_telegram_adapter.py @@ -1814,20 +1814,18 @@ async def test_telegram_media_group_max_wait_is_a_hard_deadline(): {}, asyncio.Queue(), ) - # Debounce is longer than max_wait so a quiet album would wait 0.5s. - # The hard cap must still flush around created_at + max_wait. - adapter.media_group_timeout = 0.5 - adapter.media_group_max_wait = 0.15 + # Debounce is much longer than max_wait so a quiet album would wait 5s. + # The hard cap must still flush around created_at + max_wait. Do not assert a + # tight wall-clock bound: overloaded CI can stretch a 50ms sleep past 0.3s. + adapter.media_group_timeout = 5.0 + adapter.media_group_max_wait = 0.05 delivered = asyncio.Event() - started_at = asyncio.get_running_loop().time() - processed_at: float | None = None processed_count = 0 async def process(media_group_id: str, entry: dict) -> None: - nonlocal processed_at, processed_count + nonlocal processed_count assert media_group_id == "album-deadline" processed_count = len(entry["items"]) - processed_at = asyncio.get_running_loop().time() delivered.set() adapter._process_media_group_entry = process @@ -1844,10 +1842,8 @@ async def process(media_group_id: str, entry: dict) -> None: entry = next(iter(adapter.media_group_cache.values())) assert entry["deadline"] <= entry["created_at"] + adapter.media_group_max_wait - await asyncio.wait_for(delivered.wait(), timeout=0.4) + await asyncio.wait_for(delivered.wait(), timeout=1.5) assert processed_count == 3 - assert processed_at is not None - assert processed_at - started_at < 0.3 assert not adapter.media_group_cache diff --git a/tests/unit/test_admission.py b/tests/unit/test_admission.py new file mode 100644 index 0000000000..e8b5b50b54 --- /dev/null +++ b/tests/unit/test_admission.py @@ -0,0 +1,398 @@ +"""Identity keys and overlay composition for session/sender admission.""" + +from types import SimpleNamespace + +import pytest + +from astrbot.core.auth.admission import ( + ConversationKind, + SenderAdmissionOverlay, + SessionAdmissionOverlay, + UnlistedPolicy, + compose_admission, + composed_llm_enabled, + conversation_kind_from_event, + platform_instance_from_event, + sender_admission_key_from_event, + sender_overlay_from_config, + session_admission_key, + session_admission_key_from_event, + session_overlay_from_config, +) +from astrbot.core.auth.models import Subject +from astrbot.core.pipeline.waking_check.stage import ( + WakingCheckStage, + build_unique_session_id, +) +from astrbot.core.platform.message_type import MessageType +from tests.unit.test_waking_check_stage import make_real_event + + +def test_session_admission_key_normalizes_components(): + assert ( + session_admission_key( + platform_instance="napcat", + conversation_kind="group", + conversation_id="room-a", + ) + == "session:napcat:group:room-a" + ) + encoded = session_admission_key( + platform_instance="napcat", + conversation_kind=ConversationKind.PRIVATE, + conversation_id="user with space", + ) + assert encoded.startswith("session:napcat:private:b64-") + + +def test_group_session_key_ignores_unique_session_session_id(): + event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + before = session_admission_key_from_event(event) + event.session_id = build_unique_session_id(event) + after = session_admission_key_from_event(event) + + assert event.session_id == "user-1_room-a" + assert before == after == "session:napcat:group:room-a" + assert "user-1" not in after + + event.session_id = event.get_sender_id() + assert session_admission_key_from_event(event) == "session:napcat:group:room-a" + + +def test_unique_session_on_or_off_yields_the_same_group_session_key(): + off_event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + on_event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + off_stage = WakingCheckStage() + on_stage = WakingCheckStage() + off_stage.unique_session = False + on_stage.unique_session = True + + off_stage._apply_unique_session(off_event) + on_stage._apply_unique_session(on_event) + + assert off_event.session_id == "room-a" + assert on_event.session_id == "user-1_room-a" + assert session_admission_key_from_event( + off_event + ) == session_admission_key_from_event(on_event) + assert session_admission_key_from_event(on_event) == "session:napcat:group:room-a" + + +def test_private_session_key_matches_subject_im_peer(): + event = make_real_event( + message_type=MessageType.FRIEND_MESSAGE, + group_id="", + session_id="user-1", + ) + session_key = session_admission_key_from_event(event) + sender_key = sender_admission_key_from_event(event) + subject = Subject.im( + platform_instance="napcat", + bot_account_id="bot", + sender_id="user-1", + ) + + assert conversation_kind_from_event(event) is ConversationKind.PRIVATE + assert session_key == "session:napcat:private:user-1" + assert sender_key == subject.id == "im:napcat:bot:user-1" + assert session_key.rsplit(":", 1)[-1] == sender_key.rsplit(":", 1)[-1] + + +def test_sender_key_reuses_attached_im_subject(): + event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + event.subject = Subject.im( + platform_instance="napcat-a", + bot_account_id="bot-2", + sender_id="user-9", + ) + assert sender_admission_key_from_event(event) == event.subject.id + + +def test_platform_instance_prefers_platform_id(): + event = SimpleNamespace( + get_platform_id=lambda: "instance-1", + get_platform_name=lambda: "napcat", + get_message_type=lambda: MessageType.GROUP_MESSAGE, + get_group_id=lambda: "room-a", + get_sender_id=lambda: "user-1", + get_self_id=lambda: "bot", + get_session_id=lambda: "room-a", + ) + assert platform_instance_from_event(event) == "instance-1" + assert session_admission_key_from_event(event) == "session:instance-1:group:room-a" + + +def test_platform_and_conversation_fallbacks(): + named = SimpleNamespace( + get_platform_name=lambda: "napcat", + get_message_type=lambda: MessageType.GROUP_MESSAGE.value, + get_group_id=lambda: "", + get_sender_id=lambda: "", + get_self_id=lambda: "", + get_session_id=lambda: "", + subject=Subject.guest("anon"), + ) + assert platform_instance_from_event(named) == "napcat" + assert conversation_kind_from_event(named) is ConversationKind.GROUP + assert session_admission_key_from_event(named) == "session:napcat:group:unknown" + assert sender_admission_key_from_event(named) == "im:napcat:default:unknown" + + unknown = SimpleNamespace( + get_platform_id=lambda: "", + get_platform_name=lambda: "", + get_message_type=lambda: MessageType.OTHER_MESSAGE, + get_group_id=lambda: "", + get_sender_id=lambda: "", + get_self_id=lambda: "", + get_session_id=lambda: "peer-9", + ) + assert platform_instance_from_event(unknown) == "unknown" + assert conversation_kind_from_event(unknown) is ConversationKind.PRIVATE + assert session_admission_key_from_event(unknown) == "session:unknown:private:peer-9" + + +@pytest.mark.parametrize( + ("config", "expected"), + [ + ({}, SessionAdmissionOverlay()), + ( + {"llm_enabled": False, "tts_enabled": True}, + SessionAdmissionOverlay(llm_enabled=False, listed=True), + ), + ( + {"session_enabled": False, "session_blocked": True, "llm_enabled": True}, + SessionAdmissionOverlay( + session_enabled=False, + session_blocked=True, + llm_enabled=True, + listed=True, + ), + ), + ("bad", SessionAdmissionOverlay()), + ({"session_blocked": "yes"}, SessionAdmissionOverlay()), + ({"session_enabled": None, "tts_enabled": True}, SessionAdmissionOverlay()), + ], +) +def test_session_overlay_from_config(config, expected): + assert session_overlay_from_config(config) == expected + + +@pytest.mark.parametrize( + ("config", "expected"), + [ + ({}, SenderAdmissionOverlay()), + ( + {"blocked": True, "llm_enabled": False}, + SenderAdmissionOverlay(blocked=True, llm_enabled=False, listed=True), + ), + ({"blocked": "yes"}, SenderAdmissionOverlay()), + (None, SenderAdmissionOverlay()), + ], +) +def test_sender_overlay_from_config(config, expected): + assert sender_overlay_from_config(config) == expected + + +@pytest.mark.parametrize( + ( + "label", + "session_config", + "sender_config", + "unlisted_sessions", + "unlisted_senders", + "admit_event", + "admit_llm", + ), + [ + ( + "unwritten_umo_and_uid", + {}, + {}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.ALLOW, + True, + True, + ), + ( + "session_open_uid_unwritten", + {"session_enabled": True}, + {}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.ALLOW, + True, + True, + ), + ( + "session_llm_off_uid_unwritten", + {"llm_enabled": False}, + {}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.ALLOW, + True, + False, + ), + ( + "session_open_sender_llm_off", + {"session_enabled": True}, + {"llm_enabled": False}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.ALLOW, + True, + False, + ), + ( + "vip_sender_overrides_session_llm_off", + {"llm_enabled": False}, + {"llm_enabled": True}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.ALLOW, + True, + True, + ), + ( + "sender_blocked_drops_event", + {"session_enabled": True}, + {"blocked": True}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.ALLOW, + False, + False, + ), + ( + "session_blocked_overrides_sender_llm", + {"session_blocked": True}, + {"llm_enabled": True}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.ALLOW, + False, + False, + ), + ( + "unlisted_sessions_deny_without_allow", + {}, + {"llm_enabled": True}, + UnlistedPolicy.DENY, + UnlistedPolicy.ALLOW, + False, + False, + ), + ( + "unlisted_sessions_deny_with_session_overlay", + {"session_enabled": True}, + {}, + UnlistedPolicy.DENY, + UnlistedPolicy.ALLOW, + True, + True, + ), + ( + "unlisted_senders_deny_without_uid_allow", + {"session_enabled": True}, + {}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.DENY, + False, + False, + ), + ( + "unlisted_senders_deny_with_uid_allow", + {}, + {"llm_enabled": True}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.DENY, + True, + True, + ), + ( + "unlisted_senders_deny_blocked_is_not_allow", + {}, + {"blocked": True}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.DENY, + False, + False, + ), + ( + "unlisted_senders_deny_listed_llm_off_still_admits_event", + {}, + {"llm_enabled": False}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.DENY, + True, + False, + ), + ( + "session_disabled_still_admits_event", + {"session_enabled": False}, + {}, + UnlistedPolicy.ALLOW, + UnlistedPolicy.ALLOW, + True, + True, + ), + ( + "invalid_overlay_values_are_unlisted", + {"session_blocked": "yes"}, + {"blocked": "yes"}, + UnlistedPolicy.DENY, + UnlistedPolicy.DENY, + False, + False, + ), + ], +) +def test_compose_admission_parent_table( + label, + session_config, + sender_config, + unlisted_sessions, + unlisted_senders, + admit_event, + admit_llm, +): + _ = label + decision = compose_admission( + session_overlay_from_config(session_config), + sender_overlay_from_config(sender_config), + unlisted_sessions=unlisted_sessions, + unlisted_senders=unlisted_senders, + ) + assert decision.admit_event is admit_event + assert decision.admit_llm is admit_llm + + +def test_composed_llm_enabled_ignores_event_drop_flags(): + session = session_overlay_from_config( + {"session_blocked": True, "llm_enabled": True} + ) + sender = sender_overlay_from_config({"llm_enabled": True}) + assert composed_llm_enabled(session, sender) is True + decision = compose_admission(session, sender) + assert decision.admit_event is False + assert decision.admit_llm is False + assert decision.session_blocked is True + + +def test_compose_admission_exposes_session_enabled_without_dropping_event(): + decision = compose_admission( + session_overlay_from_config({"session_enabled": False}), + sender_overlay_from_config({}), + ) + assert decision.admit_event is True + assert decision.session_enabled is False diff --git a/tests/unit/test_session_llm_manager.py b/tests/unit/test_session_llm_manager.py new file mode 100644 index 0000000000..3a5ae61cdf --- /dev/null +++ b/tests/unit/test_session_llm_manager.py @@ -0,0 +1,197 @@ +"""SessionServiceManager admission queries preserve empty-UID behavior.""" + +from typing import Any + +import pytest + +from astrbot.core.auth.admission import sender_admission_key_from_event +from astrbot.core.platform.message_type import MessageType +from astrbot.core.star.session_llm_manager import SessionServiceManager +from tests.unit.test_waking_check_stage import make_real_event + + +class _Preferences: + def __init__(self) -> None: + self.values: dict[tuple[str, str, str], Any] = {} + self.gets: list[tuple[str, str, str]] = [] + + async def get_async( + self, + scope: str, + scope_id: str, + key: str, + default: Any = None, + ) -> Any: + self.gets.append((scope, scope_id, key)) + return self.values.get((scope, scope_id, key), default) + + async def put_async( + self, + scope: str, + scope_id: str, + key: str, + value: Any, + ) -> None: + self.values[(scope, scope_id, key)] = value + + +def _manager(preferences: _Preferences | None = None) -> SessionServiceManager: + return SessionServiceManager(preferences or _Preferences()) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("umo_config", "enabled", "blocked", "llm"), + [ + ({}, True, False, True), + ({"session_enabled": True, "llm_enabled": True}, True, False, True), + ({"session_enabled": False}, False, False, True), + ({"session_blocked": True}, True, True, True), + ({"llm_enabled": False}, True, False, False), + ({"session_blocked": "yes", "llm_enabled": "no"}, True, False, True), + ], +) +async def test_empty_uid_matches_current_session_switches( + umo_config, enabled, blocked, llm +): + event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + preferences = _Preferences() + if umo_config: + preferences.values[ + ("umo", event.unified_msg_origin, "session_service_config") + ] = umo_config + manager = _manager(preferences) + + assert await manager.is_session_enabled(event.unified_msg_origin) is enabled + assert await manager.is_session_blocked(event.unified_msg_origin) is blocked + assert await manager.should_process_llm_request(event) is llm + assert await manager.is_llm_enabled_for_session(event.unified_msg_origin) is llm + assert await manager.is_sender_blocked(event) is False + assert ( + "sender", + sender_admission_key_from_event(event), + "session_service_config", + ) in ((scope, scope_id, key) for scope, scope_id, key in preferences.gets) + + +@pytest.mark.asyncio +async def test_sender_scope_is_readable_and_vip_overrides_session_llm(): + event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + sender_key = sender_admission_key_from_event(event) + preferences = _Preferences() + preferences.values[("umo", event.unified_msg_origin, "session_service_config")] = { + "llm_enabled": False + } + preferences.values[("sender", sender_key, "session_service_config")] = { + "llm_enabled": True + } + manager = _manager(preferences) + + assert await manager.is_llm_enabled_for_session(event.unified_msg_origin) is False + assert await manager.should_process_llm_request(event) is True + assert await manager.is_sender_blocked(event) is False + + +@pytest.mark.asyncio +async def test_sender_blocked_overlay_does_not_change_empty_session_llm(): + event = make_real_event( + message_type=MessageType.FRIEND_MESSAGE, + group_id="", + session_id="user-1", + ) + preferences = _Preferences() + preferences.values[ + ("sender", sender_admission_key_from_event(event), "session_service_config") + ] = {"blocked": True} + manager = _manager(preferences) + + assert await manager.is_sender_blocked(event) is True + assert await manager.should_process_llm_request(event) is True + assert await manager.is_session_enabled(event.unified_msg_origin) is True + + +@pytest.mark.asyncio +async def test_non_dict_service_config_is_treated_as_empty(): + event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + preferences = _Preferences() + preferences.values[("umo", event.unified_msg_origin, "session_service_config")] = ( + "bad" + ) + preferences.values[ + ("sender", sender_admission_key_from_event(event), "session_service_config") + ] = ["blocked"] + manager = _manager(preferences) + + assert await manager.is_session_enabled(event.unified_msg_origin) is True + assert await manager.is_session_blocked(event.unified_msg_origin) is False + assert await manager.should_process_llm_request(event) is True + assert await manager.is_sender_blocked(event) is False + + +@pytest.mark.asyncio +async def test_setters_replace_non_dict_service_config(): + preferences = _Preferences() + preferences.values[("umo", "sid", "session_service_config")] = "bad" + manager = _manager(preferences) + + await manager.set_llm_status_for_session("sid", False) + assert preferences.values[("umo", "sid", "session_service_config")] == { + "llm_enabled": False + } + + preferences.values[("umo", "sid", "session_service_config")] = ["blocked"] + await manager.set_session_blocked("sid", True) + assert preferences.values[("umo", "sid", "session_service_config")] == { + "session_blocked": True + } + + preferences.values[("umo", "sid", "session_service_config")] = "bad" + await manager.set_tts_status_for_session("sid", False) + assert preferences.values[("umo", "sid", "session_service_config")] == { + "tts_enabled": False + } + + +@pytest.mark.asyncio +async def test_tts_non_dict_and_invalid_values_default_enabled(): + event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + preferences = _Preferences() + preferences.values[("umo", event.unified_msg_origin, "session_service_config")] = ( + "bad" + ) + manager = _manager(preferences) + + assert await manager.is_tts_enabled_for_session(event.unified_msg_origin) is True + assert await manager.should_process_tts_request(event) is True + + preferences.values[("umo", event.unified_msg_origin, "session_service_config")] = { + "tts_enabled": "no", + "llm_enabled": False, + } + assert await manager.is_tts_enabled_for_session(event.unified_msg_origin) is True + assert await manager.is_llm_enabled_for_session(event.unified_msg_origin) is False + + await manager.set_tts_status_for_session(event.unified_msg_origin, False) + assert preferences.values[ + ("umo", event.unified_msg_origin, "session_service_config") + ] == { + "tts_enabled": False, + "llm_enabled": False, + } + assert await manager.is_tts_enabled_for_session(event.unified_msg_origin) is False diff --git a/tests/unit/test_waking_check_stage.py b/tests/unit/test_waking_check_stage.py index 75e8f68525..8135dec6c0 100644 --- a/tests/unit/test_waking_check_stage.py +++ b/tests/unit/test_waking_check_stage.py @@ -1143,3 +1143,31 @@ def test_qq_official_unique_session_keeps_group_id(): ) assert waking.build_unique_session_id(event) == "user-1_group-1" + + +def test_unique_session_does_not_change_session_admission_key(): + from astrbot.core.auth.admission import session_admission_key_from_event + + off_event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + on_event = make_real_event( + message_type=MessageType.GROUP_MESSAGE, + group_id="room-a", + session_id="room-a", + ) + off_stage = waking.WakingCheckStage() + on_stage = waking.WakingCheckStage() + off_stage.unique_session = False + on_stage.unique_session = True + off_stage._apply_unique_session(off_event) + on_stage._apply_unique_session(on_event) + + assert off_event.session_id == "room-a" + assert on_event.session_id == "user-1_room-a" + assert session_admission_key_from_event(off_event) == ( + session_admission_key_from_event(on_event) + ) + assert session_admission_key_from_event(on_event) == "session:napcat:group:room-a" diff --git a/tests/unit/test_whitelist_check_stage.py b/tests/unit/test_whitelist_check_stage.py index 7fd28ccdca..df02a73b01 100644 --- a/tests/unit/test_whitelist_check_stage.py +++ b/tests/unit/test_whitelist_check_stage.py @@ -5,6 +5,7 @@ from astrbot.core.auth.models import AuthContext, Resource, Subject from astrbot.core.pipeline.whitelist_check.stage import WhitelistCheckStage +from astrbot.core.platform.message_type import MessageType @pytest.mark.asyncio @@ -40,3 +41,47 @@ async def test_can_bypass_authorizes_provider_manage_against_instance_resource() Resource.instance("default"), context, ) + + +class _WhitelistEvent: + def __init__(self, *, umo: str, group_id: str) -> None: + self.unified_msg_origin = umo + self._group_id = group_id + self.stopped = False + + def get_platform_name(self) -> str: + return "napcat" + + def get_message_type(self): + return MessageType.GROUP_MESSAGE + + def get_group_id(self) -> str: + return self._group_id + + def stop_event(self) -> None: + self.stopped = True + + +@pytest.mark.asyncio +async def test_whitelist_still_matches_umo_or_group_id(): + stage = WhitelistCheckStage() + stage.enable_whitelist_check = True + stage.whitelist = ["room-a"] + stage.wl_ignore_admin_on_group = False + stage.wl_ignore_admin_on_friend = False + stage.wl_log = False + + allowed = _WhitelistEvent( + umo="napcat:GroupMessage:user-1_room-a", + group_id="room-a", + ) + denied = _WhitelistEvent( + umo="napcat:GroupMessage:user-1_room-b", + group_id="room-b", + ) + + await stage.process(allowed) + await stage.process(denied) + + assert allowed.stopped is False + assert denied.stopped is True