reduce code complex

This commit is contained in:
MuXinCG
2026-02-15 18:01:53 +08:00
parent e672f71c12
commit 68181fae49
6 changed files with 213 additions and 145 deletions
+100 -52
View File
@@ -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)
+2
View File
@@ -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:
+5 -70
View File
@@ -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
+30 -11
View File
@@ -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 ─────────────────────────────────
+9
View File
@@ -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] 心跳维持启动...
+67 -12
View File
@@ -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):