From aac36fd96f1ff7c67e4b8ca9d85048668befccde Mon Sep 17 00:00:00 2001 From: teknium1 <127238744+teknium1@users.noreply.github.com> Date: Sat, 12 Sep 2026 19:49:52 -0700 Subject: [PATCH] refactor(gateway): text-batch flush lives on BasePlatformAdapter; shield + cancel-race fixes reach all 8 adapters MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Eight adapters (discord, telegram, wecom, matrix, whatsapp, simplex, feishu, weixin) each kept a copy of the delayed text-batch flush that base._enqueue_text_event schedules. Two correctness fixes had landed in single copies only: Discord's asyncio.shield around the dispatch (#12444 — a late chunk cancelling the flush task aborted the in-flight agent turn) and WeCom/Weixin's synchronous task-identity check before the pop (a superseded task waking late popped the event and the successor found nothing). The other adapters carried both bugs latent. The base now owns `_flush_text_batch` with both fixes, plus the `_pending_text_batches` / `_pending_text_batch_tasks` dicts and `_SPLIT_THRESHOLD` / delay attrs (defaults; adapters set their own). Platform policy goes through three small hooks instead of a copied body: `_text_batch_delay_for(pending)` (Telegram's fast/short tiers, WeCom's attachment-only wait), `_pop_text_batch(key)` (Feishu's side count table) and `_dispatch_text_batch(event)` (Feishu's per-chat lock). Telegram keeps its `_flush_buffered` body because its contract differs on purpose — a cancel after the pop must hold-and-re-raise so teardown can stop a flush — and gains the identity check there. Matrix's `_split_threshold` is renamed to the shared `_SPLIT_THRESHOLD`; SimpleX exposes its single delay through the shared attr names. `helpers.TextBatchAggregator` (zero users) sits inside the revert-scheduled PLUGIN-COMPAT block and is left for that revert. --- gateway/platforms/base.py | 59 +++++++++++- gateway/platforms/weixin.py | 17 ---- plugins/platforms/discord/adapter.py | 32 ------- plugins/platforms/feishu/adapter.py | 18 +--- plugins/platforms/matrix/adapter.py | 23 +---- plugins/platforms/simplex/adapter.py | 20 +--- plugins/platforms/telegram/adapter.py | 36 ++++--- plugins/platforms/wecom/adapter.py | 29 +----- plugins/platforms/whatsapp/adapter.py | 15 --- tests/gateway/test_matrix_message_length.py | 4 +- .../test_text_batch_flush_invariants.py | 95 +++++++++++++++++++ tests/gateway/test_text_batching.py | 10 +- 12 files changed, 193 insertions(+), 165 deletions(-) create mode 100644 tests/gateway/test_text_batch_flush_invariants.py diff --git a/gateway/platforms/base.py b/gateway/platforms/base.py index 23337e02f9..add5890e5f 100644 --- a/gateway/platforms/base.py +++ b/gateway/platforms/base.py @@ -1796,6 +1796,8 @@ class BasePlatformAdapter(ABC): # could drop a newer guard. self._active_sessions: Dict[str, asyncio.Event] = {} self._pending_messages: Dict[str, MessageEvent] = {} + self._pending_text_batches: Dict[str, MessageEvent] = {} + self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} self._session_tasks: Dict[str, asyncio.Task] = {} # Legacy env knob; the runner syncs the busy_input_mode value after construction. # Default "interrupt" so a pre-sync read never silently queues. @@ -2235,8 +2237,15 @@ class BasePlatformAdapter(ABC): return None return resolved if isinstance(resolved, str) and resolved.strip() else None - # ── Inbound text batching: subclasses supply ``_pending_text_batches`` / - # ``_pending_text_batch_tasks`` dicts and ``_flush_text_batch(key)``. + # ── Inbound text batching. Chat clients split one long message into several inbound + # chunks; ``_enqueue_text_event`` merges chunks per session key and ``_flush_text_batch`` + # dispatches after a quiet period (longer when the last chunk sits near the platform's + # split point, i.e. a continuation is almost certain). Adapters set the delay attrs and + # ``_SPLIT_THRESHOLD``; ``_text_batch_delay_for`` / ``_pop_text_batch`` / + # ``_dispatch_text_batch`` are the override seams for platform-specific policy. + _SPLIT_THRESHOLD: int = 4000 + _text_batch_delay_seconds: float = 0.0 + _text_batch_split_delay_seconds: float = 0.0 def _event_session_key(self, event: "MessageEvent") -> str: """Adapter-level session key for ``event``, profile-namespaced like the agent run.""" @@ -2271,6 +2280,52 @@ class BasePlatformAdapter(ABC): prior_task.cancel() self._pending_text_batch_tasks[key] = asyncio.create_task(self._flush_text_batch(key)) + def _text_batch_delay_for(self, pending: Optional["MessageEvent"]) -> float: + """Quiet period before ``pending`` is dispatched; near-split chunks wait longer.""" + last_len = getattr(pending, "_last_chunk_len", 0) if pending is not None else 0 + return self._text_batch_split_delay_seconds if last_len >= self._SPLIT_THRESHOLD else self._text_batch_delay_seconds + + def _pop_text_batch(self, key: str) -> Optional["MessageEvent"]: + """Remove and return the pending batch for ``key`` (adapters with side tables override).""" + return self._pending_text_batches.pop(key, None) + + async def _dispatch_text_batch(self, event: "MessageEvent") -> None: + """Hand a flushed batch to the pipeline (adapters with per-chat guards override).""" + await self.handle_message(event) + + async def _flush_text_batch_now(self, key: str) -> None: + """Dispatch the pending batch for ``key`` immediately (no quiet period).""" + event = self._pop_text_batch(key) + if event is not None: + await self._dispatch_text_batch(event) + + async def _flush_text_batch(self, key: str) -> None: + """Wait for the quiet period, then dispatch the batch for ``key``. + + Two races share this body. (1) ``_enqueue_text_event`` cancels the prior flush task + on each new chunk; when ``Task.cancel()`` lands after ``sleep()`` already completed, + CancelledError is delivered at the *next* await — after a superseded task would have + popped the event, so the successor finds nothing and the message is lost. The identity + check therefore runs synchronously between the sleep and the pop. (2) A cancel that + lands while the dispatch is in flight would abort the agent turn (#12444), so the + dispatch is shielded and the outer CancelledError swallowed.""" + current_task = asyncio.current_task() + try: + await asyncio.sleep(self._text_batch_delay_for(self._pending_text_batches.get(key))) + owner = self._pending_text_batch_tasks.get(key) + if owner is not None and owner is not current_task: + return + event = self._pop_text_batch(key) + if event is None: + return + logger.info("[%s] Flushing text batch %s (%d chars)", self.name, key, len(event.text or "")) + await asyncio.shield(self._dispatch_text_batch(event)) + except asyncio.CancelledError: + pass + finally: + if self._pending_text_batch_tasks.get(key) is current_task: + self._pending_text_batch_tasks.pop(key, None) + def _history_media_paths_for_session(self, session_key: str) -> Optional[set]: """Return media paths already delivered in prior turns of this session (MEDIA: tags / image_generate payloads), so an echoed old tag isn't re-sent.""" diff --git a/gateway/platforms/weixin.py b/gateway/platforms/weixin.py index c2fc011e87..aeda977dbd 100644 --- a/gateway/platforms/weixin.py +++ b/gateway/platforms/weixin.py @@ -731,8 +731,6 @@ class WeixinAdapter(BasePlatformAdapter): # trigger a separate agent run. 3s / 5s (after a ~2048-char split chunk) suit iLink's cadence. self._text_batch_delay_seconds = self._coerce_float_extra("text_batch_delay_seconds", 3.0) self._text_batch_split_delay_seconds = self._coerce_float_extra("text_batch_split_delay_seconds", 5.0) - self._pending_text_batches: Dict[str, MessageEvent] = {} - self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} persisted = load_weixin_account(hermes_home, self._account_id) if self._account_id and not self._token else None if persisted: self._token = str(persisted.get("token") or "").strip() @@ -945,21 +943,6 @@ class WeixinAdapter(BasePlatformAdapter): event.source, group_sessions_per_user=self.config.extra.get("group_sessions_per_user", True), thread_sessions_per_user=self.config.extra.get("thread_sessions_per_user", False), profile=event.source.profile) - async def _flush_text_batch(self, key: str) -> None: - current_task = asyncio.current_task() - try: - pending = self._pending_text_batches.get(key) - split = (getattr(pending, "_last_chunk_len", 0) if pending else 0) >= self._SPLIT_THRESHOLD - await asyncio.sleep(self._text_batch_split_delay_seconds if split else self._text_batch_delay_seconds) - if self._pending_text_batch_tasks.get(key) is not current_task: - return - event = self._pending_text_batches.pop(key, None) - if event: - await self.handle_message(event) - finally: - if self._pending_text_batch_tasks.get(key) is current_task: - self._pending_text_batch_tasks.pop(key, None) - async def _collect_media(self, item: Dict[str, Any], media_paths: List[str], media_types: List[str]) -> None: spec = _INBOUND_MEDIA.get(item.get("type")) path, mime = await self._download_media(item, spec) if spec else (None, "") diff --git a/plugins/platforms/discord/adapter.py b/plugins/platforms/discord/adapter.py index b2b309b8f2..15279c467b 100644 --- a/plugins/platforms/discord/adapter.py +++ b/plugins/platforms/discord/adapter.py @@ -1034,8 +1034,6 @@ class DiscordAdapter(DiscordMediaMixin, BasePlatformAdapter): # Text batching: merge rapid successive messages (Telegram-style) self._text_batch_delay_seconds = env_float("HERMES_DISCORD_TEXT_BATCH_DELAY_SECONDS", 0.6) self._text_batch_split_delay_seconds = env_float("HERMES_DISCORD_TEXT_BATCH_SPLIT_DELAY_SECONDS", 2.0) - self._pending_text_batches: Dict[str, MessageEvent] = {} - self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} self._voice_text_channels: Dict[int, int] = {} # guild_id -> text_channel_id self._voice_sources: Dict[int, Dict[str, Any]] = {} # guild_id -> linked text channel source metadata self._voice_timeout_tasks: Dict[int, asyncio.Task] = {} # guild_id -> timeout task @@ -5868,36 +5866,6 @@ class DiscordAdapter(DiscordMediaMixin, BasePlatformAdapter): await self.handle_message(event) return True - # ------------------------------------------------------------------ - # Text message aggregation (handles Discord client-side splits) - # ------------------------------------------------------------------ - - async def _flush_text_batch(self, key: str) -> None: - """Wait for the quiet period then dispatch; longer delay when the chunk is - near Discord's 2000-char split point (continuation almost certain).""" - current_task = asyncio.current_task() - try: - pending = self._pending_text_batches.get(key) - last_len = getattr(pending, "_last_chunk_len", 0) if pending else 0 - if last_len >= self._SPLIT_THRESHOLD: - delay = self._text_batch_split_delay_seconds - else: - delay = self._text_batch_delay_seconds - await asyncio.sleep(delay) - event = self._pending_text_batches.pop(key, None) - if not event: - return - logger.info("[Discord] Flushing text batch %s (%d chars)", key, len(event.text or "")) - # Shield the dispatch: _enqueue_text_event cancels the prior flush task on each new chunk; - # without the shield CancelledError would abort the in-flight agent turn. - await asyncio.shield(self.handle_message(event)) - except asyncio.CancelledError: - # Cancel landed before the pop; shielded handle_message unaffected. - pass - finally: - if self._pending_text_batch_tasks.get(key) is current_task: - self._pending_text_batch_tasks.pop(key, None) - # --------------------------------------------------------------------------- # Discord UI Components (outside the adapter class) diff --git a/plugins/platforms/feishu/adapter.py b/plugins/platforms/feishu/adapter.py index 9c02560c0c..f6381ae94e 100644 --- a/plugins/platforms/feishu/adapter.py +++ b/plugins/platforms/feishu/adapter.py @@ -2866,21 +2866,11 @@ class FeishuAdapter(BasePlatformAdapter): prior_task.cancel() task_map[key] = asyncio.create_task(flush_fn(key)) - async def _flush_text_batch(self, key: str) -> None: - """Flush after the quiet period; wait longer when the last chunk sits near Feishu's ~4096-char split.""" - pending = self._pending_text_batches.get(key) - last_len = getattr(pending, "_last_chunk_len", 0) if pending else 0 - near_split = last_len >= self._SPLIT_THRESHOLD # a continuation chunk is almost certain - delay = self._text_batch_split_delay_seconds if near_split else self._text_batch_delay_seconds - await self._delayed_flush(self._pending_text_batch_tasks, key, delay, self._flush_text_batch_now) - - async def _flush_text_batch_now(self, key: str) -> None: - """Dispatch the current text batch immediately.""" - event = self._pending_text_batches.pop(key, None) + def _pop_text_batch(self, key: str) -> Optional[MessageEvent]: self._pending_text_batch_counts.pop(key, None) - if not event: - return - logger.info("[Feishu] Flushing text batch %s (%d chars)", key, len(event.text or "")) + return self._pending_text_batches.pop(key, None) + + async def _dispatch_text_batch(self, event: MessageEvent) -> None: await self._handle_message_with_guards(event) # --- Message content extraction and resource download --- diff --git a/plugins/platforms/matrix/adapter.py b/plugins/platforms/matrix/adapter.py index 935981b969..7cb681ac62 100644 --- a/plugins/platforms/matrix/adapter.py +++ b/plugins/platforms/matrix/adapter.py @@ -799,7 +799,7 @@ class MatrixAdapter(BasePlatformAdapter): typed_command_prefix = "!" # clients reserve typed "/" for local commands; "!command" always reaches Hermes # Class-level defaults keep object.__new__-built test instances working. max_message_length = DEFAULT_MAX_MESSAGE_LENGTH - _split_threshold = DEFAULT_MAX_MESSAGE_LENGTH - 100 + _SPLIT_THRESHOLD = DEFAULT_MAX_MESSAGE_LENGTH - 100 def _resolve_store_dir(self) -> Path: """Pin the crypto-store dir to the active profile (connect() runs inside the profile @@ -816,7 +816,7 @@ class MatrixAdapter(BasePlatformAdapter): self.max_message_length = _resolve_max_message_length(config) self.MAX_MESSAGE_LENGTH = self.max_message_length # mirrors other adapters for tooling # A chunk near the outbound limit almost certainly has a continuation. - self._split_threshold = max(100, self.max_message_length - 100) + self._SPLIT_THRESHOLD = max(100, self.max_message_length - 100) # Homeserver/user_id/device_id go through the same scoped reader as the token/password: # under multiplex os.environ holds the DEFAULT profile's identity, and pairing it with a # secondary's credential sends that credential to the wrong homeserver (or reuses the @@ -875,8 +875,6 @@ class MatrixAdapter(BasePlatformAdapter): # Text batching merges client-side splits (~4000 chars) of one long message. self._text_batch_delay_seconds = float(os.getenv("HERMES_MATRIX_TEXT_BATCH_DELAY_SECONDS", "0.6")) self._text_batch_split_delay_seconds = float(os.getenv("HERMES_MATRIX_TEXT_BATCH_SPLIT_DELAY_SECONDS", "2.0")) - self._pending_text_batches: Dict[str, MessageEvent] = {} - self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} self._approval_reaction_map = { "✅": "once", "🌀": "session", "♾️": "always", "♾": "always", "\u267e\ufe0f": "always", "\u267e": "always", "❌": "deny", "❎": "deny"} @@ -2481,23 +2479,6 @@ class MatrixAdapter(BasePlatformAdapter): except Exception as exc: logger.debug("Matrix: failed to redact model picker reaction %s: %s", emoji, exc) - async def _flush_text_batch(self, key: str) -> None: - """Wait for the quiet period then dispatch the aggregated text.""" - current_task = asyncio.current_task() - try: - pending = self._pending_text_batches.get(key) - last_len = getattr(pending, "_last_chunk_len", 0) if pending else 0 - near_split = last_len >= self._split_threshold - await asyncio.sleep(self._text_batch_split_delay_seconds if near_split else self._text_batch_delay_seconds) - event = self._pending_text_batches.pop(key, None) - if not event: - return - logger.info("[Matrix] Flushing text batch %s (%d chars)", key, len(event.text or "")) - await self.handle_message(event) - finally: - if self._pending_text_batch_tasks.get(key) is current_task: - self._pending_text_batch_tasks.pop(key, None) - def _background_read_receipt(self, room_id: str, event_id: str) -> None: async def _send() -> None: diff --git a/plugins/platforms/simplex/adapter.py b/plugins/platforms/simplex/adapter.py index 3bf900f093..7c4817a8d9 100644 --- a/plugins/platforms/simplex/adapter.py +++ b/plugins/platforms/simplex/adapter.py @@ -122,10 +122,9 @@ class SimplexAdapter(BasePlatformAdapter): self._pending_file_transfers: Dict[int, dict] = {} # awaiting rcvFileComplete, by fileId self._pending_responses: Dict[str, asyncio.Future] = {} # awaited command replies self._corr_counter = 0 - # Text batching state consumed by BasePlatformAdapter._enqueue_text_event. - self._text_batch_delay = float(os.getenv("HERMES_SIMPLEX_TEXT_BATCH_DELAY", "0.8")) - self._pending_text_batches: Dict[str, MessageEvent] = {} - self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} + # SimpleX has no client-side split, so the split delay equals the plain one. + self._text_batch_delay_seconds = float(os.getenv("HERMES_SIMPLEX_TEXT_BATCH_DELAY", "0.8")) + self._text_batch_split_delay_seconds = self._text_batch_delay_seconds logger.info( "SimpleX adapter initialized: url=%s auto_accept=%s groups=%s", self.ws_url, self.auto_accept, "enabled" if self.group_allow_from else "disabled") @@ -381,19 +380,6 @@ class SimplexAdapter(BasePlatformAdapter): def _text_batch_key(self, event: MessageEvent) -> str: return f"{event.source.platform.value}:{event.source.chat_id}" - async def _flush_text_batch(self, key: str) -> None: - """Wait for the quiet period then dispatch the aggregated text.""" - current_task = asyncio.current_task() - try: - await asyncio.sleep(self._text_batch_delay) - event = self._pending_text_batches.pop(key, None) - if event: - logger.info("[SimpleX] Flushing text batch %s (%d chars)", key, len(event.text or "")) - await self.handle_message(event) - finally: - if self._pending_text_batch_tasks.get(key) is current_task: - self._pending_text_batch_tasks.pop(key, None) - def _make_corr_id(self) -> str: """Mint a correlation ID and remember it for echo-filtering; the set is bounded by ``_max_pending_corr`` (overflow evicted in one sweep).""" diff --git a/plugins/platforms/telegram/adapter.py b/plugins/platforms/telegram/adapter.py index 659dadaa96..03edcbdd55 100644 --- a/plugins/platforms/telegram/adapter.py +++ b/plugins/platforms/telegram/adapter.py @@ -456,8 +456,6 @@ class TelegramAdapter(BasePlatformAdapter): "HERMES_TELEGRAM_TEXT_BATCH_DELAY_SECONDS", 0.3, min_value=0.08, max_value=2.0) self._text_batch_split_delay_seconds = self._env_float_clamped( "HERMES_TELEGRAM_TEXT_BATCH_SPLIT_DELAY_SECONDS", 1.0, min_value=self._text_batch_delay_seconds, max_value=4.0) - self._pending_text_batches: Dict[str, MessageEvent] = {} - self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} self._drop_delayed_deliveries = False # Held across disconnect: PTB advances the offset before our drop-guard runs, so Telegram won't # redeliver — dropping is permanent loss (see _hold_inbound_event). @@ -5810,6 +5808,11 @@ class TelegramAdapter(BasePlatformAdapter): event = None try: await asyncio.sleep(delay) + # Superseded flush (a newer chunk re-armed the timer while our sleep was already done): + # CancelledError only lands at the next await, so check synchronously before the pop. + owner = tasks.get(key) + if owner is not None and owner is not current_task: + return event = pending.pop(key, None) if not event: return @@ -5829,23 +5832,26 @@ class TelegramAdapter(BasePlatformAdapter): if tasks.get(key) is current_task: tasks.pop(key, None) - async def _flush_text_batch(self, key: str) -> None: - """Wait for the quiet period then dispatch the aggregated text.""" - # Adaptive delay: near-split-point last chunk → long delay (continuation almost certain); - # short/medium totals → capped fast delays; else configured cap (all min()'d with the operator cap). - pending = self._pending_text_batches.get(key) + def _text_batch_delay_for(self, pending: Optional[MessageEvent]) -> float: + """Adaptive delay: near-split-point last chunk → long delay (continuation almost certain); + short/medium totals → capped fast delays; else configured cap (all min()'d with the operator cap).""" last_len = getattr(pending, "_last_chunk_len", 0) if pending else 0 total_len = len(getattr(pending, "text", "") or "") if pending else 0 if last_len >= self._SPLIT_THRESHOLD: - delay = self._text_batch_split_delay_seconds - elif total_len <= self._TEXT_BATCH_FAST_LEN: - delay = min(self._text_batch_delay_seconds, self._TEXT_BATCH_FAST_DELAY_S) - elif total_len <= self._TEXT_BATCH_SHORT_LEN: - delay = min(self._text_batch_delay_seconds, self._TEXT_BATCH_SHORT_DELAY_S) - else: - delay = self._text_batch_delay_seconds + return self._text_batch_split_delay_seconds + if total_len <= self._TEXT_BATCH_FAST_LEN: + return min(self._text_batch_delay_seconds, self._TEXT_BATCH_FAST_DELAY_S) + if total_len <= self._TEXT_BATCH_SHORT_LEN: + return min(self._text_batch_delay_seconds, self._TEXT_BATCH_SHORT_DELAY_S) + return self._text_batch_delay_seconds + + async def _flush_text_batch(self, key: str) -> None: + """Telegram keeps its own flush body: a cancel after the pop must HOLD the event and re-raise + (PTB already acked the update; the hold queue redispatches after reconnect) rather than shield + the dispatch — teardown must be able to stop a flush from reaching a torn-down session.""" await self._flush_buffered( - self._pending_text_batches, self._pending_text_batch_tasks, key, delay, "text", + self._pending_text_batches, self._pending_text_batch_tasks, key, + self._text_batch_delay_for(self._pending_text_batches.get(key)), "text", lambda ev: logger.info("[Telegram] Flushing text batch %s (%d chars)", key, len(ev.text or ""))) # -- Photo batching -- diff --git a/plugins/platforms/wecom/adapter.py b/plugins/platforms/wecom/adapter.py index 66eaa6d2fe..4116da4289 100644 --- a/plugins/platforms/wecom/adapter.py +++ b/plugins/platforms/wecom/adapter.py @@ -149,8 +149,6 @@ class WeComAdapter(WeComStreamMixin, WeComMediaMixin, ChatSendQueueMixin, BasePl self._text_batch_delay_seconds = env_float("HERMES_WECOM_TEXT_BATCH_DELAY_SECONDS", 0.6) self._text_batch_split_delay_seconds = env_float("HERMES_WECOM_TEXT_BATCH_SPLIT_DELAY_SECONDS", 2.0) self._attachment_text_merge_delay_seconds = _extra_float("attachment_text_merge_delay_seconds", 0.8) - self._pending_text_batches: Dict[str, MessageEvent] = {} - self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} # Stream keep-alive config (see streaming.py STREAM_* constants). self._stream_safe_duration_seconds = _extra_float("stream_safe_duration_seconds", STREAM_SAFE_DURATION_SECONDS) self._stream_keepalive_enabled = bool(extra.get("stream_keepalive_enabled", STREAM_KEEPALIVE_ENABLED_DEFAULT)) @@ -503,29 +501,10 @@ class WeComAdapter(WeComStreamMixin, WeComMediaMixin, ChatSendQueueMixin, BasePl existing.reply_to_text = event.reply_to_text existing.reply_to_message_id = event.reply_to_message_id - async def _flush_text_batch(self, key: str) -> None: - current_task = asyncio.current_task() - try: - pending = self._pending_text_batches.get(key) - if pending and pending.media_urls and not (pending.text or "").strip(): - delay = self._attachment_text_merge_delay_seconds # attachment-only: wait for text - elif pending and getattr(pending, "_last_chunk_len", 0) >= self._SPLIT_THRESHOLD: - delay = self._text_batch_split_delay_seconds # continuation almost certain - else: - delay = self._text_batch_delay_seconds - await asyncio.sleep(delay) - # Cancel-delivery race: CancelledError lands at the NEXT await, so this identity check - # must stay synchronous (no await between it and the pop). - if self._pending_text_batch_tasks.get(key) is not current_task: - return - event = self._pending_text_batches.pop(key, None) - if not event: - return - logger.info("[WeCom] Flushing batch %s (%d chars, %d media)", key, len(event.text or ""), len(event.media_urls or [])) - await self.handle_message(event) - finally: - if self._pending_text_batch_tasks.get(key) is current_task: - self._pending_text_batch_tasks.pop(key, None) + def _text_batch_delay_for(self, pending: Optional[MessageEvent]) -> float: + if pending is not None and pending.media_urls and not (pending.text or "").strip(): + return self._attachment_text_merge_delay_seconds # attachment-only: wait for the text frame + return super()._text_batch_delay_for(pending) @staticmethod def _extract_text(body: Dict[str, Any]) -> Tuple[str, Optional[str]]: diff --git a/plugins/platforms/whatsapp/adapter.py b/plugins/platforms/whatsapp/adapter.py index 09f407554d..1a918b5a0a 100644 --- a/plugins/platforms/whatsapp/adapter.py +++ b/plugins/platforms/whatsapp/adapter.py @@ -283,8 +283,6 @@ class WhatsAppAdapter(WhatsAppBehaviorMixin, BasePlatformAdapter): # Text debounce batching: rapid bursts (forwards, paste-splits) would otherwise each trigger a separate agent turn. self._text_batch_delay_seconds = self._coerce_float_extra("text_batch_delay_seconds", 5.0) self._text_batch_split_delay_seconds = self._coerce_float_extra("text_batch_split_delay_seconds", 10.0) - self._pending_text_batches: Dict[str, MessageEvent] = {} - self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} def _coerce_float_extra(self, key: str, default: float) -> float: """Read a float from ``config.extra``; NaN/Inf/negative/unparseable → ``default`` (fed to asyncio.sleep).""" @@ -732,19 +730,6 @@ class WhatsAppAdapter(WhatsAppBehaviorMixin, BasePlatformAdapter): _SPLIT_THRESHOLD = 6000 # WhatsApp supports ~65K chars; generous threshold - async def _flush_text_batch(self, key: str) -> None: - current_task = asyncio.current_task() - try: - pending = self._pending_text_batches.get(key) - last_len = getattr(pending, "_last_chunk_len", 0) if pending else 0 - await asyncio.sleep(self._text_batch_split_delay_seconds if last_len >= self._SPLIT_THRESHOLD else self._text_batch_delay_seconds) - event = self._pending_text_batches.pop(key, None) - if event: - await self.handle_message(event) - finally: - if self._pending_text_batch_tasks.get(key) is current_task: - self._pending_text_batch_tasks.pop(key, None) - @staticmethod def _classify_bridge_message(data: Dict[str, Any]) -> MessageType: media_type = str(data.get("mediaType", "") or "") diff --git a/tests/gateway/test_matrix_message_length.py b/tests/gateway/test_matrix_message_length.py index f3d8a2996e..a10a538f54 100644 --- a/tests/gateway/test_matrix_message_length.py +++ b/tests/gateway/test_matrix_message_length.py @@ -27,12 +27,12 @@ class TestMatrixMaxMessageLength: def test_default_limit_is_16000(self): adapter = _make_adapter() assert adapter.max_message_length == 16000 - assert adapter._split_threshold == 15900 + assert adapter._SPLIT_THRESHOLD == 15900 def test_extra_override(self): adapter = _make_adapter(max_message_length=12000) assert adapter.max_message_length == 12000 - assert adapter._split_threshold == 11900 + assert adapter._SPLIT_THRESHOLD == 11900 def test_env_override(self, monkeypatch): monkeypatch.setenv("MATRIX_MAX_MESSAGE_LENGTH", "20000") diff --git a/tests/gateway/test_text_batch_flush_invariants.py b/tests/gateway/test_text_batch_flush_invariants.py new file mode 100644 index 0000000000..7ae1699cf4 --- /dev/null +++ b/tests/gateway/test_text_batch_flush_invariants.py @@ -0,0 +1,95 @@ +"""Invariants for BasePlatformAdapter's inbound text-batch flush. + +Every adapter that batches inbound text used to carry its own ``_flush_text_batch``; only Discord +shielded the dispatch (#12444) and only WeCom/Weixin checked task identity before the pop. Both +fixes now live in the base flush; these tests pin them against the base class and sweep every +adapter that still overrides the flush (Telegram keeps a hold-queue variant) for the same contract. +""" + +import asyncio +from typing import Any, Dict + +import pytest + +from gateway.config import Platform, PlatformConfig +from gateway.platforms.base import BasePlatformAdapter +from gateway.platforms.event import MessageEvent, MessageType +from gateway.session import SessionSource + + +class _Adapter(BasePlatformAdapter): + def __init__(self): + super().__init__(PlatformConfig(enabled=True), Platform.TELEGRAM) + self._text_batch_delay_seconds = 0.0 + self._text_batch_split_delay_seconds = 0.0 + self.dispatched: list = [] + self.entered = asyncio.Event() + self.release = asyncio.Event() + + async def connect(self, *, is_reconnect: bool = False) -> bool: + return True + + async def disconnect(self) -> None: + pass + + async def send(self, *a: Any, **k: Any) -> None: + pass + + async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: + return {} + + async def handle_message(self, event: MessageEvent) -> None: + self.entered.set() + await self.release.wait() # a real agent turn awaits here + self.dispatched.append(event) + + +def _event(text: str) -> MessageEvent: + return MessageEvent(text=text, message_type=MessageType.TEXT, + source=SessionSource(platform=Platform.TELEGRAM, chat_id="c", chat_type="dm")) + + +@pytest.mark.asyncio +async def test_cancel_mid_dispatch_does_not_abort_the_turn(): + """A cancel landing while handle_message is in flight (a late chunk re-arming the timer) must not + abort the dispatch: the batch is delivered exactly once and the task exits clean.""" + adapter = _Adapter() + adapter._pending_text_batches["k"] = _event("hello") + task = asyncio.create_task(adapter._flush_text_batch("k")) + adapter._pending_text_batch_tasks["k"] = task + await adapter.entered.wait() + task.cancel() + adapter.release.set() + await task # no CancelledError escapes + assert [e.text for e in adapter.dispatched] == ["hello"] + assert adapter._pending_text_batch_tasks == {} + + +@pytest.mark.asyncio +async def test_superseded_flush_leaves_batch_for_its_successor(): + """When a newer flush task owns the key by the time the old one wakes, the old one must neither + pop nor dispatch — otherwise the successor finds an empty batch and the message is lost.""" + adapter = _Adapter() + adapter.release.set() + event = _event("hello") + adapter._pending_text_batches["k"] = event + stale = asyncio.create_task(adapter._flush_text_batch("k")) + adapter._pending_text_batch_tasks["k"] = stale + successor = asyncio.create_task(asyncio.sleep(60)) + adapter._pending_text_batch_tasks["k"] = successor # re-armed before `stale` ran + await stale + successor.cancel() + assert adapter.dispatched == [] + assert adapter._pending_text_batches.get("k") is event + assert adapter._pending_text_batch_tasks.get("k") is successor + + +@pytest.mark.asyncio +async def test_enqueue_then_flush_delivers_merged_text_once(): + adapter = _Adapter() + adapter.release.set() + adapter._enqueue_text_event(_event("part one")) + adapter._enqueue_text_event(_event("part two")) + await asyncio.gather(*adapter._pending_text_batch_tasks.values()) + assert [e.text for e in adapter.dispatched] == ["part one\npart two"] + assert adapter._pending_text_batches == {} and adapter._pending_text_batch_tasks == {} diff --git a/tests/gateway/test_text_batching.py b/tests/gateway/test_text_batching.py index 63504c0b74..e8be4fe27c 100644 --- a/tests/gateway/test_text_batching.py +++ b/tests/gateway/test_text_batching.py @@ -45,7 +45,7 @@ def _make_discord_adapter(): config = PlatformConfig(enabled=True, token="test-token") adapter = object.__new__(DiscordAdapter) - adapter._platform = Platform.DISCORD + adapter._platform = adapter.platform = Platform.DISCORD adapter.config = config adapter._pending_text_batches = {} adapter._pending_text_batch_tasks = {} @@ -105,7 +105,7 @@ def _make_matrix_adapter(): config = PlatformConfig(enabled=True, token="test-token") adapter = object.__new__(MatrixAdapter) - adapter._platform = Platform.MATRIX + adapter._platform = adapter.platform = Platform.MATRIX adapter.config = config adapter._pending_text_batches = {} adapter._pending_text_batch_tasks = {} @@ -159,7 +159,7 @@ def _make_wecom_adapter(): config = PlatformConfig(enabled=True, token="test-token") adapter = object.__new__(WeComAdapter) - adapter._platform = Platform.WECOM + adapter._platform = adapter.platform = Platform.WECOM adapter.config = config adapter._pending_text_batches = {} adapter._pending_text_batch_tasks = {} @@ -213,7 +213,7 @@ def _make_telegram_adapter(): config = PlatformConfig(enabled=True, token="test-token") adapter = object.__new__(TelegramAdapter) - adapter._platform = Platform.TELEGRAM + adapter._platform = adapter.platform = Platform.TELEGRAM adapter.config = config adapter._pending_text_batches = {} adapter._pending_text_batch_tasks = {} @@ -247,7 +247,7 @@ def _make_feishu_adapter(): config = PlatformConfig(enabled=True, token="test-token") adapter = object.__new__(FeishuAdapter) - adapter._platform = Platform.FEISHU + adapter._platform = adapter.platform = Platform.FEISHU adapter.config = config batch_state = FeishuBatchState() adapter._pending_text_batches = batch_state.events