refactor(gateway): text-batch flush lives on BasePlatformAdapter; shield + cancel-race fixes reach all 8 adapters
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.
This commit is contained in:
@@ -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."""
|
||||
|
||||
@@ -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, "")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 ---
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)."""
|
||||
|
||||
@@ -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 --
|
||||
|
||||
@@ -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]]:
|
||||
|
||||
@@ -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 "")
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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 == {}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user