From 68181fae49ef0bab3f59f3982cc58d573d740af2 Mon Sep 17 00:00:00 2001 From: MuXinCG <202322130196@mail.sdu.edu.cn> Date: Sun, 15 Feb 2026 18:01:53 +0800 Subject: [PATCH] reduce code complex --- EvoScientist/channels/base.py | 152 ++++++++++++------ EvoScientist/channels/bus/events.py | 2 + EvoScientist/channels/channel_manager.py | 75 +-------- EvoScientist/channels/imessage/channel_rpc.py | 41 +++-- botpy.log | 9 ++ tests/test_channel_comprehensive.py | 79 +++++++-- 6 files changed, 213 insertions(+), 145 deletions(-) diff --git a/EvoScientist/channels/base.py b/EvoScientist/channels/base.py index edd475b..6a77f46 100644 --- a/EvoScientist/channels/base.py +++ b/EvoScientist/channels/base.py @@ -9,7 +9,6 @@ import asyncio import dataclasses import logging import re -import time from collections import defaultdict from collections.abc import Awaitable, Callable as CallableABC from dataclasses import dataclass, field @@ -291,10 +290,6 @@ class Channel(ChannelPlugin, ABC): self._message_ids: dict[str, str] = {} self._debounce_tasks: dict[str, asyncio.Task] = {} - # Deduplication - from .middleware import DedupCache - self._dedup = DedupCache() - # Mention gating: "always" | "group" | "off" self.require_mention: str = getattr(config, "require_mention", "group") @@ -312,9 +307,56 @@ class Channel(ChannelPlugin, ABC): # Per-chat send locks to prevent message reordering self._send_locks: dict[str, asyncio.Lock] = defaultdict(asyncio.Lock) - # Group history buffer for context injection - from .middleware import GroupHistoryBuffer - self._group_history = GroupHistoryBuffer() + # Build inbound middleware pipeline + self._inbound_middlewares = self._build_inbound_middlewares() + + def _build_inbound_middlewares(self) -> list: + """Build the inbound middleware chain from config and capabilities. + + Middleware order: + 1. DedupMiddleware — drop duplicates early + 2. AllowListMiddleware — enforce sender/channel restrictions + 3. PairingMiddleware — handle DM pairing (if applicable) + 4. GroupHistoryMiddleware — buffer/inject group history + 5. MentionGatingMiddleware — filter by mention policy + """ + from .middleware import ( + DedupMiddleware, AllowListMiddleware, + PairingMiddleware, GroupHistoryMiddleware, MentionGatingMiddleware, + ) + middlewares = [] + middlewares.append(DedupMiddleware()) + # AllowList + allowed_senders = getattr(self.config, "allowed_senders", None) + allowed_channels = getattr(self.config, "allowed_channels", None) + if allowed_senders and not isinstance(allowed_senders, set): + allowed_senders = set(allowed_senders) + if allowed_channels and not isinstance(allowed_channels, set): + allowed_channels = set(allowed_channels) + middlewares.append(AllowListMiddleware( + allowed_senders=allowed_senders, + allowed_channels=allowed_channels, + dm_policy=self.dm_policy, + )) + # Pairing + if self.dm_policy == "pairing": + async def _send_pair(chat_id, text): + await self._send_chunk(chat_id, text, text, None, {}) + middlewares.append(PairingMiddleware( + channel_name=self.name, + send_response_fn=_send_pair, + dm_policy=self.dm_policy, + )) + # GroupHistory + if self.capabilities.groups: + middlewares.append(GroupHistoryMiddleware()) + # MentionGating + if self.capabilities.mentions: + middlewares.append(MentionGatingMiddleware( + require_mention=self.require_mention, + strip_fn=self._strip_mention, + )) + return middlewares @abstractmethod async def start(self) -> None: @@ -721,9 +763,37 @@ class Channel(ChannelPlugin, ABC): # ── Inbound message pipeline ────────────────────────────────────── + def _raw_to_inbound(self, raw: RawIncoming) -> InboundMessage | None: + """Convert a RawIncoming to InboundMessage (pure transformation, no filtering). + + Merges text + annotations into content, sets metadata. + Returns None only if there is no content and no media. + """ + parts = [] + if raw.text: + parts.append(raw.text) + parts.extend(raw.content_annotations) + content = "\n".join(p for p in parts if p) + if not content and not raw.media_files: + return None + meta = dict(raw.metadata) + meta.setdefault("chat_id", raw.chat_id) + return InboundMessage( + channel=self.name, sender_id=raw.sender_id, chat_id=raw.chat_id, + content=content or "[media only]", timestamp=raw.timestamp, + message_id=raw.message_id, media=raw.media_files, metadata=meta, + is_group=raw.is_group, was_mentioned=raw.was_mentioned, + ) + def _build_inbound(self, raw: RawIncoming) -> InboundMessage | None: """Build an ``InboundMessage`` from raw platform data. + .. deprecated:: + This method is superseded by the middleware pipeline in + ``_enqueue_raw()``. It is kept for backward compatibility + with tests that call it directly. New code should use + ``_enqueue_raw()`` instead. + Performs mention gating, allow-list checks, and merges text + annotations into a single content string. Returns ``None`` if the message should be dropped (not mentioned, sender not @@ -792,44 +862,28 @@ class Channel(ChannelPlugin, ABC): ) async def _enqueue_raw(self, raw: RawIncoming) -> None: - """Build an InboundMessage from *raw* and put it on the queue. + """Run *raw* through the inbound middleware pipeline, convert to + InboundMessage, and put it on the queue. Convenience method for subclass ``_on_message`` handlers. - Stores non-mentioned group messages in history buffer and injects - context when the bot is mentioned. """ - from .middleware import HistoryEntry - - # Store all group messages in history buffer BEFORE mention gating - if raw.is_group: - ts = raw.timestamp.timestamp() if hasattr(raw.timestamp, 'timestamp') else time.time() - if not raw.was_mentioned: - # Store for context, _build_inbound will return None via _should_process - self._group_history.add(raw.chat_id, HistoryEntry( - sender_id=raw.sender_id, - text=raw.text, - timestamp=ts, - message_id=raw.message_id, - )) - else: - # Inject context before the current message - context = self._group_history.format_context(raw.chat_id) - if context: - raw = dataclasses.replace( - raw, - text=context + "\n\n[Current message - respond to this]\n" + raw.text, - ) - self._group_history.clear(raw.chat_id) - - msg = self._build_inbound(raw) - if msg is not None: - # Fire-and-forget ACK reaction - if raw.message_id: - try: - await self._send_ack_reaction(raw.chat_id, raw.message_id) - except Exception: - pass - await self._queue.put(msg) + context: dict = {"channel": self} + current: RawIncoming | None = raw + for mw in self._inbound_middlewares: + if current is None: + return + current = await mw.process_inbound(current, context) + if current is None: + return + msg = self._raw_to_inbound(current) + if msg is None: + return + if raw.message_id: + try: + await self._send_ack_reaction(raw.chat_id, raw.message_id) + except Exception: + pass + await self._queue.put(msg) # ── Bus integration ────────────────────────────────────────────── @@ -838,22 +892,16 @@ class Channel(ChannelPlugin, ABC): self._bus = bus async def queue_message(self, msg: InboundMessage) -> None: - """Buffer *msg* with debounce + dedup, then publish to bus.""" + """Buffer *msg* with debounce, then publish to bus.""" sender = msg.sender_id - # Deduplication - mid = msg.message_id - if mid and self._dedup.is_duplicate(mid): - _logger.debug(f"Dedup: skipping duplicate message {mid}") - return - if sender not in self._message_buffers: self._message_buffers[sender] = [] self._message_metadata[sender] = msg.metadata self._message_media[sender] = [] self._message_buffers[sender].append(msg.content) - if mid: - self._message_ids[sender] = mid + if msg.message_id: + self._message_ids[sender] = msg.message_id if msg.media: self._message_media[sender].extend(msg.media) diff --git a/EvoScientist/channels/bus/events.py b/EvoScientist/channels/bus/events.py index 8e94727..411006f 100644 --- a/EvoScientist/channels/bus/events.py +++ b/EvoScientist/channels/bus/events.py @@ -21,6 +21,8 @@ class InboundMessage: message_id: str = "" media: list[str] = field(default_factory=list) metadata: dict[str, Any] = field(default_factory=dict) + is_group: bool = False + was_mentioned: bool = True @property def sender(self) -> str: diff --git a/EvoScientist/channels/channel_manager.py b/EvoScientist/channels/channel_manager.py index bc2ffb0..2a1dd7f 100644 --- a/EvoScientist/channels/channel_manager.py +++ b/EvoScientist/channels/channel_manager.py @@ -248,17 +248,6 @@ class AccountManager: from .middleware import ( InboundMiddleware, OutboundMiddlewareBase, - DedupMiddleware, - AllowListMiddleware, - MentionGatingMiddleware, - GroupHistoryMiddleware, - DebounceMiddleware, - AckReactionMiddleware, - FormattingMiddleware, - ChunkingMiddleware, - RetryMiddleware, - TypingMiddleware, - PairingMiddleware, ) @@ -314,67 +303,16 @@ class OutboundPipeline: return current -def build_inbound_pipeline( - plugin: ChannelPlugin, - config: Any, -) -> InboundPipeline: - """Auto-assemble inbound pipeline based on plugin capabilities. - - Middleware order: - 1. DedupMiddleware — drop duplicates early - 2. AllowListMiddleware — enforce sender/channel restrictions - 3. PairingMiddleware — handle DM pairing (if applicable) - 4. GroupHistoryMiddleware — buffer/inject group history - 5. MentionGatingMiddleware — filter by mention policy - """ - caps = plugin.capabilities - middlewares: list[InboundMiddleware] = [] - - middlewares.append(DedupMiddleware()) - - allowed_senders = getattr(config, "allowed_senders", None) - allowed_channels = getattr(config, "allowed_channels", None) - dm_policy = getattr(config, "dm_policy", "allowlist") - if allowed_senders: - allowed_senders = set(allowed_senders) if not isinstance(allowed_senders, set) else allowed_senders - if allowed_channels: - allowed_channels = set(allowed_channels) if not isinstance(allowed_channels, set) else allowed_channels - middlewares.append(AllowListMiddleware( - allowed_senders=allowed_senders, - allowed_channels=allowed_channels, - dm_policy=dm_policy, - )) - - if plugin.pairing is not None: - middlewares.append(PairingMiddleware( - channel_name=plugin.id, - dm_policy=dm_policy, - )) - - if caps.groups: - middlewares.append(GroupHistoryMiddleware()) - - if caps.mentions: - strip_fn = None - if plugin.mentions is not None: - strip_fn = lambda text, _adapter=plugin.mentions: _adapter.strip_mentions(text, {}) # noqa: E731 - require_mention = getattr(config, "require_mention", "group") - middlewares.append(MentionGatingMiddleware( - require_mention=require_mention, - strip_fn=strip_fn, - )) - - return InboundPipeline(plugin, middlewares) - - def build_outbound_pipeline( plugin: ChannelPlugin, config: Any, ) -> OutboundPipeline: - """Auto-assemble outbound pipeline based on plugin capabilities.""" - caps = plugin.capabilities + """Auto-assemble outbound pipeline based on plugin capabilities. + + FormattingMiddleware has been removed — Channel.send() handles + formatting + chunking via _format_chunk() / _prepare_chunks(). + """ middlewares: list[OutboundMiddlewareBase] = [] - middlewares.append(FormattingMiddleware(caps)) return OutboundPipeline(plugin, middlewares) @@ -672,7 +610,6 @@ class ChannelManager: self._health_providers: dict[str, Callable[[], dict]] = {} self._account_manager = AccountManager() # Pipelines (built during registration) - self._inbound_pipelines: dict[str, InboundPipeline] = {} self._outbound_pipelines: dict[str, OutboundPipeline] = {} # Shared webhook self._shared_webhook_port = shared_webhook_port @@ -742,7 +679,6 @@ class ChannelManager: if channel.config_adapter is not None: self._account_manager.register_plugin(channel) if config is not None: - self._inbound_pipelines[name] = build_inbound_pipeline(channel, config) self._outbound_pipelines[name] = build_outbound_pipeline(channel, config) logger.info(f"Registered channel: {name} (slots: {channel.filled_slots()})") return channel @@ -1086,7 +1022,6 @@ class ChannelManager: "total_successes": health.total_successes, }, "plugin_slots": channel.filled_slots(), - "has_inbound_pipeline": name in self._inbound_pipelines, "has_outbound_pipeline": name in self._outbound_pipelines, } return result diff --git a/EvoScientist/channels/imessage/channel_rpc.py b/EvoScientist/channels/imessage/channel_rpc.py index 2b2f094..ebf32d6 100644 --- a/EvoScientist/channels/imessage/channel_rpc.py +++ b/EvoScientist/channels/imessage/channel_rpc.py @@ -24,6 +24,25 @@ from .targets import ( logger = logging.getLogger(__name__) +class _IMessageAllowListMiddleware: + """Custom allow-list middleware for iMessage's rich sender filtering. + + Supports chat_id/chat_guid matching, wildcard, and normalized + phone/email matching — logic that the generic AllowListMiddleware + does not cover. + """ + + def __init__(self, channel: 'IMessageChannelRpc'): + self._channel = channel + + async def process_inbound(self, raw, context): + chat_id = raw.metadata.get("chat_id") + chat_guid = raw.metadata.get("chat_guid") + if not self._channel._is_sender_allowed(raw.sender_id, chat_id, chat_guid): + return None + return raw + + @dataclass class IMessageConfig(BaseChannelConfig): """Configuration for iMessage channel.""" @@ -55,18 +74,18 @@ class IMessageChannelRpc(Channel): # ── Pipeline overrides ──────────────────────────────────────── - def is_allowed(self, sender: str) -> bool: - # Rich filtering is handled in _build_inbound via _is_sender_allowed - return True + def _build_inbound_middlewares(self): + """Use iMessage-specific allow-list middleware. - def _build_inbound(self, raw): - """Override to apply iMessage's rich sender filtering.""" - chat_id = raw.metadata.get("chat_id") - chat_guid = raw.metadata.get("chat_guid") - if not self._is_sender_allowed(raw.sender_id, chat_id, chat_guid): - logger.debug(f"Ignoring message from {raw.sender_id}") - return None - return super()._build_inbound(raw) + iMessage doesn't need MentionGating (always sets was_mentioned=True). + """ + from ..middleware import DedupMiddleware, GroupHistoryMiddleware + middlewares = [] + middlewares.append(DedupMiddleware()) + middlewares.append(_IMessageAllowListMiddleware(self)) + if self.capabilities.groups: + middlewares.append(GroupHistoryMiddleware()) + return middlewares # ── Incoming message handling ───────────────────────────────── diff --git a/botpy.log b/botpy.log index b43c983..eeddecd 100644 --- a/botpy.log +++ b/botpy.log @@ -16,3 +16,12 @@ 2026-02-15 13:57:49,073 [INFO] (gateway.py:142)ws_identify [botpy] 鉴权中... 2026-02-15 13:57:49,294 [INFO] (gateway.py:85)on_message [botpy] 机器人「EvoSci-测试中」启动成功! 2026-02-15 13:57:49,295 [INFO] (gateway.py:223)_send_heart [botpy] 心跳维持启动... +2026-02-15 18:01:03,252 [INFO] (client.py:162)_bot_login [botpy] 登录机器人账号中... +2026-02-15 18:01:03,471 [INFO] (robot.py:65)update_access_token [botpy] access_token expires_in 2348 +2026-02-15 18:01:03,935 [INFO] (client.py:181)_bot_init [botpy] 程序启动... +2026-02-15 18:01:03,935 [INFO] (connection.py:60)multi_run [botpy] 最大并发连接数: 1, 启动会话数: 1 +2026-02-15 18:01:03,935 [INFO] (client.py:242)bot_connect [botpy] 会话启动中... +2026-02-15 18:01:03,935 [INFO] (gateway.py:115)ws_connect [botpy] 启动中... +2026-02-15 18:01:04,118 [INFO] (gateway.py:142)ws_identify [botpy] 鉴权中... +2026-02-15 18:01:04,298 [INFO] (gateway.py:85)on_message [botpy] 机器人「EvoSci-测试中」启动成功! +2026-02-15 18:01:04,299 [INFO] (gateway.py:223)_send_heart [botpy] 心跳维持启动... diff --git a/tests/test_channel_comprehensive.py b/tests/test_channel_comprehensive.py index ee45860..6b0e9d9 100644 --- a/tests/test_channel_comprehensive.py +++ b/tests/test_channel_comprehensive.py @@ -64,6 +64,7 @@ class _FakeConfig: allowed_channels: list | None = None proxy: str | None = None require_mention: str = "group" + dm_policy: str = "allowlist" class StubChannel(Channel): @@ -618,6 +619,64 @@ class TestChannelBuildInbound: assert msg.metadata["extra"] == "data" +class TestInboundPipeline: + """Tests for the new middleware-based inbound pipeline in _enqueue_raw().""" + + def test_pipeline_dedup(self): + """Duplicate messages are dropped by the pipeline.""" + async def _test(): + ch = StubChannel() + raw = RawIncoming(sender_id="u1", chat_id="c1", text="hello", message_id="m1") + await ch._enqueue_raw(raw) + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 1 + _run(_test()) + + def test_pipeline_allowlist_blocks(self): + """Non-allowed senders are blocked by the pipeline.""" + async def _test(): + cfg = _FakeConfig(allowed_senders=["alice"]) + ch = StubChannel(cfg) + raw = RawIncoming(sender_id="eve", chat_id="c1", text="hack") + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 0 + _run(_test()) + + def test_pipeline_allowlist_passes(self): + """Allowed senders pass through the pipeline.""" + async def _test(): + cfg = _FakeConfig(allowed_senders=["alice"]) + ch = StubChannel(cfg) + raw = RawIncoming(sender_id="alice", chat_id="c1", text="hello") + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 1 + _run(_test()) + + def test_pipeline_channel_allowlist_blocks(self): + """Non-allowed channels are blocked by the pipeline.""" + async def _test(): + cfg = _FakeConfig(allowed_channels=["c1"]) + ch = StubChannel(cfg) + raw = RawIncoming(sender_id="u1", chat_id="c2", text="hello") + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 0 + _run(_test()) + + def test_pipeline_inbound_has_is_group(self): + """InboundMessage carries is_group and was_mentioned from RawIncoming.""" + async def _test(): + ch = StubChannel() + raw = RawIncoming( + sender_id="u1", chat_id="c1", text="hello", + is_group=True, was_mentioned=True, + ) + await ch._enqueue_raw(raw) + msg = await ch._queue.get() + assert msg.is_group is True + assert msg.was_mentioned is True + _run(_test()) + + class TestChannelDebounce: def test_single_message_processed(self): @@ -670,23 +729,19 @@ class TestChannelDebounce: _run(_test()) def test_dedup_skips_duplicate(self): + """Dedup is now handled in _enqueue_raw pipeline, not queue_message.""" async def _test(): - bus = MessageBus() ch = StubChannel() - ch.set_bus(bus) - ch.initial_debounce = 0.05 - msg = InboundMessage( - channel="stub", sender_id="u1", chat_id="c1", - content="hello", message_id="m1", - metadata={"chat_id": "c1"}, + raw = RawIncoming( + sender_id="u1", chat_id="c1", text="hello", + message_id="m1", ) - await ch.queue_message(msg) - await ch.queue_message(msg) # duplicate - await asyncio.sleep(0.2) + await ch._enqueue_raw(raw) + await ch._enqueue_raw(raw) # duplicate - # Only one should be processed (dedup catches second) - assert bus.inbound.qsize() == 1 + # Only one should be enqueued (dedup catches second) + assert ch._queue.qsize() == 1 _run(_test()) def test_debounce_metadata_from_first_message(self):