reduce code complex
This commit is contained in:
+100
-52
@@ -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)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 ─────────────────────────────────
|
||||
|
||||
|
||||
@@ -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] 心跳维持启动...
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user