diff --git a/plugins/platforms/matrix/adapter.py b/plugins/platforms/matrix/adapter.py index 391f2fe531..a313d0c443 100644 --- a/plugins/platforms/matrix/adapter.py +++ b/plugins/platforms/matrix/adapter.py @@ -398,7 +398,7 @@ 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, is_voice_event +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 @@ -2028,17 +2028,15 @@ class MatrixAdapter(BasePlatformAdapter): async def _resolve_message_context( self, room_id: str, sender: str, event_id: str, body: str, source_content: dict, - relates_to: dict) -> Optional[tuple]: + relates_to: dict, mention_claimed: bool = False) -> Optional[tuple]: """Shared mention/thread/DM gating. Returns (body, is_dm, chat_type, thread_id, - display_name, source) or None when the message should be dropped.""" + display_name, source) or None when the message should be dropped. ``mention_claimed`` + marks a parked voice claimed by the sender's follow-up bare @mention.""" identity = await self._resolve_room_identity(room_id) is_dm = await self._is_dm_room(room_id) chat_type = "dm" if is_dm else "group" thread_id = relates_to.get("event_id") if relates_to.get("rel_type") == "m.thread" else None - formatted_body = source_content.get("formatted_body") - mentions_block = source_content.get("m.mentions") or {} # MSC3952: authoritative signal - mention_user_ids = mentions_block.get("user_ids") if isinstance(mentions_block, dict) else None - is_mentioned = self._is_bot_mentioned(body, formatted_body, mention_user_ids) + is_mentioned = mention_claimed or self._content_mentions_bot(body, source_content) if not is_dm: # Whitelist first: non-listed rooms are dropped even when @mentioned (DMs exempt). if self._allowed_rooms and room_id not in self._allowed_rooms: @@ -2048,8 +2046,7 @@ class MatrixAdapter(BasePlatformAdapter): is_free_room = room_id in self._free_rooms in_bot_thread = bool(thread_id and thread_id in self._threads) if self._require_mention and not is_free_room and not in_bot_thread: - if (not is_mentioned and not body.startswith("/") - and not self._parked_voices.consume_claim(event_id)): + if not is_mentioned and not body.startswith("/"): if is_voice_event(source_content): # a bare @mention may follow (Element X) self._parked_voices.park(room_id, sender, event_id, source_content, relates_to) logger.debug( @@ -2144,15 +2141,15 @@ class MatrixAdapter(BasePlatformAdapter): body = source_content.get("body", "") or "" if not body: return - mentions = source_content.get("m.mentions") or {} - if (self._require_mention and not self._strip_mention(body).strip() and self._is_bot_mentioned( - body, source_content.get("formatted_body"), - mentions.get("user_ids") if isinstance(mentions, dict) else None)): + # 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) + and not self._strip_mention(body).strip() and self._content_mentions_bot(body, source_content)): 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") + room_id, sender, voice_id, event_ts, voice_content, voice_relates, "m.audio", + mention_claimed=True) return msg_event = await self._build_inbound_event( room_id, sender, event_id, _normalize_matrix_bang_command(body), source_content, relates_to) @@ -2165,7 +2162,7 @@ class MatrixAdapter(BasePlatformAdapter): async def _handle_media_message( self, room_id: str, sender: str, event_id: str, event_ts: float, source_content: dict, - relates_to: dict, msgtype: str) -> None: + relates_to: dict, msgtype: str, mention_claimed: bool = False) -> None: body = source_content.get("body", "") or "" url = source_content.get("url", "") if url and not str(url).startswith("mxc://"): @@ -2194,7 +2191,8 @@ 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) + ctx = await self._resolve_message_context( + room_id, sender, event_id, body, source_content, relates_to, mention_claimed=mention_claimed) if ctx is None: return # Cache locally so downstream tools get a real file path. @@ -2222,7 +2220,7 @@ class MatrixAdapter(BasePlatformAdapter): if msgtype == "m.image": return MessageType.PHOTO, event_mimetype or "image/png", False if msgtype == "m.audio": - is_voice = source_content.get("org.matrix.msc3245.voice") is not None + is_voice = has_voice_marker(source_content) return (MessageType.VOICE if is_voice else MessageType.AUDIO), event_mimetype or "audio/ogg", is_voice if msgtype == "m.video": return MessageType.VIDEO, event_mimetype or "video/mp4", False @@ -2896,6 +2894,12 @@ class MatrixAdapter(BasePlatformAdapter): return True return bool(formatted_body and self._user_id and f"matrix.to/#/{self._user_id}" in formatted_body) + def _content_mentions_bot(self, body: str, content: dict) -> bool: + """``_is_bot_mentioned`` fed from an event's content (MSC3952 ``m.mentions`` is authoritative).""" + mentions = content.get("m.mentions") or {} + return self._is_bot_mentioned( + body, content.get("formatted_body"), mentions.get("user_ids") if isinstance(mentions, dict) else None) + def _user_localpart(self) -> str: """``@bot:server`` -> ``bot``; empty when the user ID has no server part.""" return self._user_id.split(":")[0].lstrip("@") if self._user_id and ":" in self._user_id else "" diff --git a/plugins/platforms/matrix/voice_mention.py b/plugins/platforms/matrix/voice_mention.py index f364ad2c1b..8fa89537d4 100644 --- a/plugins/platforms/matrix/voice_mention.py +++ b/plugins/platforms/matrix/voice_mention.py @@ -9,13 +9,18 @@ the window claims it. Unmentioned voices are never downloaded or transcribed whi from __future__ import annotations import time -from typing import Dict, Optional, Set, Tuple +from typing import Dict, Optional, Tuple CLAIM_WINDOW_SECONDS = 120.0 +def has_voice_marker(content: dict) -> bool: + """The single MSC3245 voice-message check (shared with the adapter's media classifier).""" + return content.get("org.matrix.msc3245.voice") is not None + + def is_voice_event(content: dict) -> bool: - return content.get("msgtype") == "m.audio" and "org.matrix.msc3245.voice" in content + return content.get("msgtype") == "m.audio" and has_voice_marker(content) class ParkedVoices: @@ -23,29 +28,25 @@ class ParkedVoices: self._window = window # (room_id, sender) -> (parked_at, event_id, content, relates_to) self._parked: Dict[Tuple[str, str], Tuple[float, str, dict, dict]] = {} - self._claimed: Set[str] = set() def _prune(self) -> None: cutoff = time.monotonic() - self._window 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 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) def claim(self, room_id: str, sender: str) -> Optional[Tuple[str, dict, dict]]: - """Pop the sender's parked voice for this room; the claimed event id then passes the - mention gate exactly once (see ``consume_claim``) instead of being re-parked.""" + """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 - self._claimed.add(event_id) return event_id, content, relates_to - - def consume_claim(self, event_id: str) -> bool: - if event_id in self._claimed: - self._claimed.discard(event_id) - return True - return False diff --git a/tests/gateway/test_matrix_mention.py b/tests/gateway/test_matrix_mention.py index ed2adfeb95..f92ba5dbde 100644 --- a/tests/gateway/test_matrix_mention.py +++ b/tests/gateway/test_matrix_mention.py @@ -243,10 +243,15 @@ async def test_bare_mention_passes_empty_string(monkeypatch): @pytest.mark.asyncio -@pytest.mark.parametrize("mention_room, claims", [("!room1:example.org", True), ("!room2:example.org", False)]) -async def test_bare_mention_claims_parked_voice_only_in_same_room(monkeypatch, mention_room, claims): +@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), +]) +async def test_bare_mention_claims_parked_voice_only_in_same_room(monkeypatch, mention_room, mention_body, claims): """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.""" + 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.""" monkeypatch.delenv("MATRIX_REQUIRE_MENTION", raising=False) monkeypatch.delenv("MATRIX_FREE_RESPONSE_ROOMS", raising=False) monkeypatch.setenv("MATRIX_AUTO_THREAD", "false") @@ -259,8 +264,9 @@ async def test_bare_mention_claims_parked_voice_only_in_same_room(monkeypatch, m 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( - "@hermes:example.org", event_id="$text", room_id=mention_room, mention_user_ids=["@hermes:example.org"])) + mention_body, event_id="$text", room_id=mention_room, mention_user_ids=["@hermes:example.org"])) 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")])