diff --git a/plugins/platforms/matrix/adapter.py b/plugins/platforms/matrix/adapter.py index a313d0c443..4e3a665540 100644 --- a/plugins/platforms/matrix/adapter.py +++ b/plugins/platforms/matrix/adapter.py @@ -77,6 +77,7 @@ from gateway.platforms.base import ( from gateway.platforms.base import transcode_to_ogg_opus from gateway.platforms.event import MessageEvent, MessageType, ProcessingOutcome from gateway.platforms.helpers import ThreadParticipationTracker +from plugins.platforms.matrix.voice_mention import ParkedVoices, has_voice_marker, is_voice_event logger = logging.getLogger(__name__) @@ -398,7 +399,6 @@ def _resolve_max_message_length(config) -> int: # identity in one crypto.db. # Store directory for E2EE keys and sync state. Mirrors the pairing-store fix (a6397c379). See #89168. from hermes_constants import get_hermes_dir as _get_hermes_dir -from plugins.platforms.matrix.voice_mention import ParkedVoices, has_voice_marker, is_voice_event _STARTUP_GRACE_SECONDS = 5 # ignore messages older than this many seconds before startup @@ -2141,15 +2141,18 @@ class MatrixAdapter(BasePlatformAdapter): body = source_content.get("body", "") or "" if not body: return - # Dict lookup first: the mention regexes only run when a voice is actually parked. - if (self._require_mention and self._parked_voices.has(room_id, sender) + # Dict lookup first: the mention regexes only run when a voice is parked or being gated + # (both only happen under require_mention). + if (self._parked_voices.pending(room_id, sender) and not self._strip_mention(body).strip() and self._content_mentions_bot(body, source_content)): + await self._parked_voices.settle(room_id, sender) # same-/sync-batch voice still gating parked = self._parked_voices.claim(room_id, sender) if parked: # answer the voice this bare mention was typed for, not an empty text voice_id, voice_content, voice_relates = parked await self._handle_media_message( room_id, sender, voice_id, event_ts, voice_content, voice_relates, "m.audio", mention_claimed=True) + self._background_read_receipt(room_id, event_id) # the claim receipted the voice return msg_event = await self._build_inbound_event( room_id, sender, event_id, _normalize_matrix_bang_command(body), source_content, relates_to) @@ -2191,8 +2194,15 @@ class MatrixAdapter(BasePlatformAdapter): msg_type, media_type, is_voice_message = self._classify_inbound_media(msgtype, event_mimetype, source_content) # Gate (require_mention / allowed rooms) BEFORE the download: an unmentioned or # non-allowlisted room must not pull media onto the host only to drop it. - ctx = await self._resolve_message_context( - room_id, sender, event_id, body, source_content, relates_to, mention_claimed=mention_claimed) + # First await: mark a parkable voice in-flight so a concurrent bare mention waits for it. + gate = (self._parked_voices.begin(room_id, sender) + if self._require_mention and not mention_claimed and is_voice_event(source_content) else None) + try: + ctx = await self._resolve_message_context( + room_id, sender, event_id, body, source_content, relates_to, mention_claimed=mention_claimed) + finally: + if gate is not None: + self._parked_voices.release(room_id, sender, gate) if ctx is None: return # Cache locally so downstream tools get a real file path. diff --git a/plugins/platforms/matrix/voice_mention.py b/plugins/platforms/matrix/voice_mention.py index 8fa89537d4..f17f313b8c 100644 --- a/plugins/platforms/matrix/voice_mention.py +++ b/plugins/platforms/matrix/voice_mention.py @@ -8,10 +8,16 @@ the window claims it. Unmentioned voices are never downloaded or transcribed whi from __future__ import annotations +import asyncio import time from typing import Dict, Optional, Tuple CLAIM_WINDOW_SECONDS = 120.0 +# How long a bare mention waits for a voice from the same /sync batch that is still being gated. +SETTLE_TIMEOUT_SECONDS = 5.0 + +# (voice event_id, content, relates_to) +ParkedVoice = Tuple[str, dict, dict] def has_voice_marker(content: dict) -> bool: @@ -24,29 +30,49 @@ def is_voice_event(content: dict) -> bool: class ParkedVoices: - def __init__(self, window: float = CLAIM_WINDOW_SECONDS) -> None: - self._window = window - # (room_id, sender) -> (parked_at, event_id, content, relates_to) - self._parked: Dict[Tuple[str, str], Tuple[float, str, dict, dict]] = {} + def __init__(self) -> None: + # (room_id, sender) -> (parked_at, parked voice) + self._parked: Dict[Tuple[str, str], Tuple[float, ParkedVoice]] = {} + # (room_id, sender) -> set once that sender's voice has been gated (parked or not). + # mautrix runs one /sync batch's events as concurrent tasks, so the voice may still be + # awaiting room identity when its bare mention is handled. + self._inflight: Dict[Tuple[str, str], asyncio.Event] = {} def _prune(self) -> None: - cutoff = time.monotonic() - self._window + cutoff = time.monotonic() - CLAIM_WINDOW_SECONDS self._parked = {k: v for k, v in self._parked.items() if v[0] >= cutoff} - def has(self, room_id: str, sender: str) -> bool: - """Cheap pre-check (may include an expired entry; ``claim`` prunes it).""" - return (room_id, sender) in self._parked + def pending(self, room_id: str, sender: str) -> bool: + """Cheap pre-check: a voice is parked (maybe expired; ``claim`` prunes) or still being gated.""" + key = (room_id, sender) + return key in self._parked or key in self._inflight + + def begin(self, room_id: str, sender: str) -> asyncio.Event: + """Mark a voice as being gated. Call before the first await; always pair with ``release``.""" + gate = self._inflight[(room_id, sender)] = asyncio.Event() + return gate + + def release(self, room_id: str, sender: str, gate: asyncio.Event) -> None: + gate.set() + if self._inflight.get((room_id, sender)) is gate: + del self._inflight[(room_id, sender)] + + async def settle(self, room_id: str, sender: str) -> None: + """Wait (bounded) for a concurrently gated voice from this sender to be parked or dropped.""" + gate = self._inflight.get((room_id, sender)) + if gate is not None: + try: + await asyncio.wait_for(gate.wait(), SETTLE_TIMEOUT_SECONDS) + except asyncio.TimeoutError: + pass def park(self, room_id: str, sender: str, event_id: str, content: dict, relates_to: dict) -> None: self._prune() - self._parked[(room_id, sender)] = (time.monotonic(), event_id, content, relates_to) + self._parked[(room_id, sender)] = (time.monotonic(), (event_id, content, relates_to)) - def claim(self, room_id: str, sender: str) -> Optional[Tuple[str, dict, dict]]: + def claim(self, room_id: str, sender: str) -> Optional[ParkedVoice]: """Pop the sender's parked voice for this room (the caller re-dispatches it with ``mention_claimed=True`` so it passes the mention gate without re-parking).""" self._prune() entry = self._parked.pop((room_id, sender), None) - if entry is None: - return None - _ts, event_id, content, relates_to = entry - return event_id, content, relates_to + return entry[1] if entry else None diff --git a/tests/gateway/test_matrix_mention.py b/tests/gateway/test_matrix_mention.py index f92ba5dbde..bc24959e2c 100644 --- a/tests/gateway/test_matrix_mention.py +++ b/tests/gateway/test_matrix_mention.py @@ -243,33 +243,51 @@ async def test_bare_mention_passes_empty_string(monkeypatch): @pytest.mark.asyncio -@pytest.mark.parametrize("mention_room, mention_body, claims", [ - ("!room1:example.org", "@hermes:example.org", True), - ("!room2:example.org", "@hermes:example.org", False), - ("!room1:example.org", "@hermes:example.org hi", False), +@pytest.mark.parametrize("mention_room, mention_body, claims, same_sync_batch", [ + ("!room1:example.org", "@hermes:example.org", True, False), + ("!room2:example.org", "@hermes:example.org", False, False), + ("!room1:example.org", "@hermes:example.org hi", False, False), + ("!room1:example.org", "@hermes:example.org", True, True), ]) -async def test_bare_mention_claims_parked_voice_only_in_same_room(monkeypatch, mention_room, mention_body, claims): +async def test_bare_mention_claims_parked_voice_only_in_same_room( + monkeypatch, mention_room, mention_body, claims, same_sync_batch): """An unmentioned MSC3245 voice (empty m.mentions) is answered by the sender's bare @mention typed right after it in the SAME room; a bare mention in another room never pulls it across, - and a mention carrying text is answered as that text.""" + and a mention carrying text is answered as that text. mautrix runs one /sync batch's events as + concurrent tasks, so the claim must also win while the voice still awaits a room-identity fetch.""" + import asyncio + monkeypatch.delenv("MATRIX_REQUIRE_MENTION", raising=False) monkeypatch.delenv("MATRIX_FREE_RESPONSE_ROOMS", raising=False) monkeypatch.setenv("MATRIX_AUTO_THREAD", "false") adapter = _make_adapter() adapter._download_and_cache_media = AsyncMock(return_value="/tmp/voice.ogg") + adapter._background_read_receipt = MagicMock() voice = _make_event("voice message", event_id="$voice") voice.content.update({"msgtype": "m.audio", "url": "mxc://example.org/v", "info": {"mimetype": "audio/ogg"}, "org.matrix.msc3245.voice": {}, "m.mentions": {}}) + mention = _make_event(mention_body, event_id="$text", room_id=mention_room, + mention_user_ids=["@hermes:example.org"]) - await adapter._on_room_message(voice) - adapter.handle_message.assert_not_awaited() - adapter._download_and_cache_media.assert_not_awaited() # parked voice is never downloaded - await adapter._on_room_message(_make_event( - mention_body, event_id="$text", room_id=mention_room, mention_user_ids=["@hermes:example.org"])) + if same_sync_batch: + resolve_identity = adapter._resolve_room_identity + + async def slow_identity(room_id): # stale 60s cache -> homeserver round-trip + await asyncio.sleep(0.01) + return await resolve_identity(room_id) + adapter._resolve_room_identity = slow_identity + await asyncio.gather(adapter._on_room_message(voice), adapter._on_room_message(mention)) + else: + await adapter._on_room_message(voice) + adapter.handle_message.assert_not_awaited() + adapter._download_and_cache_media.assert_not_awaited() # parked voice is never downloaded + await adapter._on_room_message(mention) dispatched = [(m.args[0].source.chat_id, m.args[0].message_id) for m in adapter.handle_message.await_args_list] assert dispatched == ([("!room1:example.org", "$voice")] if claims else [(mention_room, "$text")]) + if claims: # the bare mention is the newest event; the read marker must reach it + adapter._background_read_receipt.assert_any_call("!room1:example.org", "$text") # ---------------------------------------------------------------------------