From f5d493aebfa647ad0023b545274f1a3d7821cf47 Mon Sep 17 00:00:00 2001 From: yu-xin-c <2182712990@qq.com> Date: Thu, 9 Jul 2026 23:04:34 +0800 Subject: [PATCH] fix(gateway): dedupe pending voice transcript echoes --- gateway/run.py | 304 +++++++++++++----- .../test_telegram_voice_v0_regressions.py | 109 +++++++ 2 files changed, 328 insertions(+), 85 deletions(-) diff --git a/gateway/run.py b/gateway/run.py index aa4c3471a2b..a2625f4e986 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -5799,7 +5799,34 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew # at the next check point. if effective_mode == "interrupt" and running_agent and running_agent is not _AGENT_PENDING_SENTINEL: try: - running_agent.interrupt(event.text) + _interrupt_text = event.text + _media_urls = getattr(event, "media_urls", None) or [] + if self._pending_event_audio_paths(event): + try: + _interrupt_text, _transcripts = await self._transcribe_pending_audio_event_once( + event, + event.text or "", + ) + _echo_meta = self._thread_metadata_for_source( + event.source, + self._reply_anchor_for_event(event), + ) + await self._echo_pending_stt_transcripts_once( + event, + adapter, + event.source, + _transcripts, + metadata=_echo_meta, + log_context="Voice-busy-interrupt", + ) + except Exception as _trans_exc: + logger.warning( + "Voice-busy-interrupt transcription failed: %s", + _trans_exc, + ) + elif not _interrupt_text and _media_urls: + _interrupt_text = _build_media_placeholder(event) + running_agent.interrupt(_interrupt_text) except Exception: pass # don't let interrupt failure block the ack @@ -10450,7 +10477,35 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew self._queue_or_replace_pending_event(_quick_key, event) return None logger.debug("PRIORITY interrupt for session %s", _quick_key) - running_agent.interrupt(event.text) + _interrupt_text = event.text + _media_urls = getattr(event, "media_urls", None) or [] + if self._pending_event_audio_paths(event): + try: + _interrupt_text, _transcripts = await self._transcribe_pending_audio_event_once( + event, + event.text or "", + ) + _echo_adapter = self._adapter_for_source(source) + _echo_meta = self._thread_metadata_for_source( + source, + self._reply_anchor_for_event(event), + ) + await self._echo_pending_stt_transcripts_once( + event, + _echo_adapter, + source, + _transcripts, + metadata=_echo_meta, + log_context="Voice-priority-interrupt", + ) + except Exception as _trans_exc: + logger.warning( + "Voice-priority-interrupt transcription failed: %s", + _trans_exc, + ) + elif not _interrupt_text and _media_urls: + _interrupt_text = _build_media_placeholder(event) + running_agent.interrupt(_interrupt_text) # NOTE: self._pending_messages was write-only (never consumed). # The actual interrupt message is delivered via adapter._pending_messages # which is read by _run_agent. Removed to prevent unbounded growth. @@ -16317,6 +16372,80 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew return prefix, successful_transcripts return user_text, successful_transcripts + def _pending_event_audio_paths(self, event) -> List[str]: + """Return audio paths that the pending-message interrupt path transcribes.""" + audio_paths: List[str] = [] + media_urls = getattr(event, "media_urls", None) or [] + media_types = getattr(event, "media_types", None) or [] + for i, path in enumerate(media_urls): + mtype = media_types[i] if i < len(media_types) else "" + is_audio = ( + mtype.startswith("audio/") + or getattr(event, "message_type", None) in (MessageType.VOICE, MessageType.AUDIO) + ) + if is_audio: + audio_paths.append(path) + return audio_paths + + async def _transcribe_pending_audio_event_once( + self, + event, + user_text: Optional[str] = None, + ) -> tuple[str | None, List[str]]: + """Transcribe a pending audio event once and cache the result on the event. + + Voice follow-ups can be inspected first by the interrupt monitor and + later consumed by the pending-drain path. Both need the same transcript, + but only one STT call and one transcript echo should happen for the + platform message. + """ + if hasattr(event, "_gateway_pending_stt_text"): + cached_text = getattr(event, "_gateway_pending_stt_text") + cached_transcripts = getattr(event, "_gateway_pending_stt_transcripts", []) or [] + return cached_text, list(cached_transcripts) + + audio_paths = self._pending_event_audio_paths(event) + if not audio_paths: + return user_text if user_text is not None else (getattr(event, "text", None) or None), [] + + text = user_text if user_text is not None else (getattr(event, "text", "") or "") + enriched_text, successful_transcripts = await self._enrich_message_with_transcription( + text, + audio_paths, + ) + setattr(event, "_gateway_pending_stt_text", enriched_text) + setattr(event, "_gateway_pending_stt_transcripts", list(successful_transcripts)) + return enriched_text, successful_transcripts + + async def _echo_pending_stt_transcripts_once( + self, + event, + adapter, + source, + transcripts: List[str], + *, + metadata=None, + log_context: str = "Transcript", + ) -> None: + """Echo pending-event STT transcripts to the chat at most once.""" + if ( + not transcripts + or not self._should_echo_stt_transcripts() + or adapter is None + or getattr(event, "_gateway_pending_stt_echo_sent", False) + ): + return + setattr(event, "_gateway_pending_stt_echo_sent", True) + for tx in transcripts: + try: + await adapter.send( + source.chat_id, + f'🎙️ "{tx}"', + metadata=metadata, + ) + except Exception as echo_exc: + logger.debug("%s echo failed (non-fatal): %s", log_context, echo_exc) + async def _dequeue_pending_with_transcription( self, adapter, @@ -16343,42 +16472,26 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew text = event.text or "" - audio_paths: List[str] = [] - media_urls = getattr(event, "media_urls", None) or [] - media_types = getattr(event, "media_types", None) or [] - for i, path in enumerate(media_urls): - mtype = media_types[i] if i < len(media_types) else "" - is_audio = ( - mtype.startswith("audio/") - or getattr(event, "message_type", None) in (MessageType.VOICE, MessageType.AUDIO) - ) - if is_audio: - audio_paths.append(path) - - if audio_paths: - enriched_text, successful_transcripts = await self._enrich_message_with_transcription( - text, audio_paths, + if self._pending_event_audio_paths(event): + enriched_text, successful_transcripts = await self._transcribe_pending_audio_event_once( + event, + text, ) # Echo raw transcripts back to the user when configured so voice # interrupts feel identical to fresh voice messages. - if successful_transcripts and self._should_echo_stt_transcripts(): - echo_adapter = self._adapter_for_source(source) - echo_meta = {"thread_id": source.thread_id} if source.thread_id else None - if echo_adapter: - for tx in successful_transcripts: - try: - await echo_adapter.send( - source.chat_id, - f'🎙️ "{tx}"', - metadata=echo_meta, - ) - except Exception as echo_exc: - logger.debug( - "Transcript echo failed (non-fatal): %s", echo_exc, - ) + echo_adapter = self._adapter_for_source(source) + echo_meta = {"thread_id": source.thread_id} if source.thread_id else None + await self._echo_pending_stt_transcripts_once( + event, + echo_adapter, + source, + successful_transcripts, + metadata=echo_meta, + ) return enriched_text or None # Non-audio fallback: preserve original _dequeue_pending_text semantics. + media_urls = getattr(event, "media_urls", None) or [] if not text and media_urls: text = _build_media_placeholder(event) return text or None @@ -20658,36 +20771,22 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew # of fresh voice messages including the # optional 🎙️ echo back to the user. _media_urls = getattr(_peek_event, "media_urls", None) or [] - _media_types = getattr(_peek_event, "media_types", None) or [] - _audio_paths = [] - for _i, _path in enumerate(_media_urls): - _mtype = _media_types[_i] if _i < len(_media_types) else "" - _is_audio = ( - _mtype.startswith("audio/") - or getattr(_peek_event, "message_type", None) in (MessageType.VOICE, MessageType.AUDIO) - ) - if _is_audio: - _audio_paths.append(_path) - if _audio_paths: + if self._pending_event_audio_paths(_peek_event): try: - _enriched, _transcripts = await self._enrich_message_with_transcription( - pending_text, _audio_paths, + _enriched, _transcripts = await self._transcribe_pending_audio_event_once( + _peek_event, + pending_text, ) pending_text = _enriched - if _transcripts and self._should_echo_stt_transcripts(): - _echo_meta = {"thread_id": source.thread_id} if source.thread_id else None - for _tx in _transcripts: - try: - await _adapter.send( - source.chat_id, - f'🎙️ "{_tx}"', - metadata=_echo_meta, - ) - except Exception as _echo_exc: - logger.debug( - "Voice-interrupt echo failed (non-fatal): %s", - _echo_exc, - ) + _echo_meta = {"thread_id": source.thread_id} if source.thread_id else None + await self._echo_pending_stt_transcripts_once( + _peek_event, + _adapter, + source, + _transcripts, + metadata=_echo_meta, + log_context="Voice-interrupt", + ) except Exception as _trans_exc: logger.warning( "Voice-interrupt transcription failed: %s", _trans_exc, @@ -20876,6 +20975,30 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew and _backup_adapter.has_pending_interrupt(session_key)): _bp_event = _backup_adapter._pending_messages.get(session_key) _bp_text = _bp_event.text if _bp_event else None + if _bp_event is not None: + _bp_media_urls = getattr(_bp_event, "media_urls", None) or [] + if self._pending_event_audio_paths(_bp_event): + try: + _bp_text, _bp_transcripts = await self._transcribe_pending_audio_event_once( + _bp_event, + _bp_text or "", + ) + _bp_echo_meta = {"thread_id": source.thread_id} if source.thread_id else None + await self._echo_pending_stt_transcripts_once( + _bp_event, + _backup_adapter, + source, + _bp_transcripts, + metadata=_bp_echo_meta, + log_context="Voice-backup-interrupt", + ) + except Exception as _bp_trans_exc: + logger.warning( + "Voice-backup-interrupt transcription failed: %s", + _bp_trans_exc, + ) + elif not _bp_text and _bp_media_urls: + _bp_text = _build_media_placeholder(_bp_event) logger.info( "Backup interrupt detected for session %s " "(monitor task state: %s)", @@ -20936,6 +21059,30 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew and _backup_adapter.has_pending_interrupt(session_key)): _bp_event = _backup_adapter._pending_messages.get(session_key) _bp_text = _bp_event.text if _bp_event else None + if _bp_event is not None: + _bp_media_urls = getattr(_bp_event, "media_urls", None) or [] + if self._pending_event_audio_paths(_bp_event): + try: + _bp_text, _bp_transcripts = await self._transcribe_pending_audio_event_once( + _bp_event, + _bp_text or "", + ) + _bp_echo_meta = {"thread_id": source.thread_id} if source.thread_id else None + await self._echo_pending_stt_transcripts_once( + _bp_event, + _backup_adapter, + source, + _bp_transcripts, + metadata=_bp_echo_meta, + log_context="Voice-backup-interrupt", + ) + except Exception as _bp_trans_exc: + logger.warning( + "Voice-backup-interrupt transcription failed: %s", + _bp_trans_exc, + ) + elif not _bp_text and _bp_media_urls: + _bp_text = _build_media_placeholder(_bp_event) logger.info( "Backup interrupt detected for session %s " "(monitor task state: %s)", @@ -21080,35 +21227,22 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew # fresh voice messages. _pending_text = pending_event.text or "" _media_urls = getattr(pending_event, "media_urls", None) or [] - _media_types = getattr(pending_event, "media_types", None) or [] - _audio_paths = [] - for _i, _path in enumerate(_media_urls): - _mtype = _media_types[_i] if _i < len(_media_types) else "" - _is_audio = ( - _mtype.startswith("audio/") - or getattr(pending_event, "message_type", None) in (MessageType.VOICE, MessageType.AUDIO) - ) - if _is_audio: - _audio_paths.append(_path) - if _audio_paths: + if self._pending_event_audio_paths(pending_event): try: - _enriched, _transcripts = await self._enrich_message_with_transcription( - _pending_text, _audio_paths, + _enriched, _transcripts = await self._transcribe_pending_audio_event_once( + pending_event, + _pending_text, ) pending = _enriched or None - if _transcripts and self._should_echo_stt_transcripts(): - _echo_meta = {"thread_id": source.thread_id} if source.thread_id else None - for _tx in _transcripts: - try: - await adapter.send( - source.chat_id, - f'🎙️ "{_tx}"', - metadata=_echo_meta, - ) - except Exception as _echo_exc: - logger.debug( - "Voice-drain echo failed (non-fatal): %s", _echo_exc, - ) + _echo_meta = {"thread_id": source.thread_id} if source.thread_id else None + await self._echo_pending_stt_transcripts_once( + pending_event, + adapter, + source, + _transcripts, + metadata=_echo_meta, + log_context="Voice-drain", + ) except Exception as _trans_exc: logger.warning( "Voice-drain transcription failed: %s", _trans_exc, diff --git a/tests/gateway/test_telegram_voice_v0_regressions.py b/tests/gateway/test_telegram_voice_v0_regressions.py index 6e2143573c4..aa36cf65476 100644 --- a/tests/gateway/test_telegram_voice_v0_regressions.py +++ b/tests/gateway/test_telegram_voice_v0_regressions.py @@ -1,6 +1,7 @@ import sys from pathlib import Path from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -10,6 +11,7 @@ if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from gateway.config import Platform +from gateway.platforms.base import MessageEvent, MessageType from plugins.platforms.telegram.adapter import TelegramAdapter from gateway.run import GatewayRunner from gateway.session import SessionSource @@ -34,6 +36,113 @@ def _runner(adapter=None): return runner +@pytest.mark.asyncio +async def test_pending_voice_interrupt_reuses_transcript_and_echo(): + adapter = SimpleNamespace(send=AsyncMock()) + runner = _runner(adapter) + source = _source() + event = MessageEvent( + text="", + message_type=MessageType.VOICE, + source=source, + media_urls=["/tmp/telegram-voice.ogg"], + media_types=["audio/ogg"], + ) + + with patch( + "tools.transcription_tools.transcribe_audio", + return_value={"success": True, "transcript": "hello once", "provider": "mock"}, + ) as mock_transcribe: + interrupt_text, interrupt_transcripts = await runner._transcribe_pending_audio_event_once( + event, + event.text, + ) + await runner._echo_pending_stt_transcripts_once( + event, + adapter, + source, + interrupt_transcripts, + ) + + drain_text, drain_transcripts = await runner._transcribe_pending_audio_event_once( + event, + event.text, + ) + await runner._echo_pending_stt_transcripts_once( + event, + adapter, + source, + drain_transcripts, + ) + + assert interrupt_text == '"hello once"' + assert drain_text == interrupt_text + assert drain_transcripts == interrupt_transcripts == ["hello once"] + mock_transcribe.assert_called_once_with("/tmp/telegram-voice.ogg") + adapter.send.assert_awaited_once_with( + "12345", + '🎙️ "hello once"', + metadata=None, + ) + + +@pytest.mark.asyncio +async def test_busy_voice_interrupt_transcribes_before_pending_drain(monkeypatch): + adapter = SimpleNamespace(send=AsyncMock(), _pending_messages={}) + runner = _runner(adapter) + runner._is_user_authorized = lambda _source: True + runner._draining = False + runner._running_agents = {} + runner._busy_input_mode = "interrupt" + runner._busy_text_mode = "interrupt" + runner._busy_ack_ts = {} + runner._queued_events = {} + runner._agent_has_active_subagents = lambda _agent: False + runner._session_has_compression_in_flight = lambda _session_key: False + session_key = "telegram:dm:12345" + agent = MagicMock() + runner._running_agents[session_key] = agent + source = _source() + event = MessageEvent( + text="", + message_type=MessageType.VOICE, + source=source, + media_urls=["/tmp/telegram-busy-voice.ogg"], + media_types=["audio/ogg"], + ) + monkeypatch.setenv("HERMES_GATEWAY_BUSY_ACK_ENABLED", "false") + + with ( + patch("tools.approval.has_blocking_approval", return_value=False), + patch( + "tools.transcription_tools.transcribe_audio", + return_value={"success": True, "transcript": "interrupt me", "provider": "mock"}, + ) as mock_transcribe, + ): + handled = await runner._handle_active_session_busy_message(event, session_key) + drain_text, drain_transcripts = await runner._transcribe_pending_audio_event_once( + adapter._pending_messages[session_key], + event.text, + ) + await runner._echo_pending_stt_transcripts_once( + adapter._pending_messages[session_key], + adapter, + source, + drain_transcripts, + ) + + assert handled is True + agent.interrupt.assert_called_once_with('"interrupt me"') + assert adapter._pending_messages[session_key] is event + assert drain_text == '"interrupt me"' + mock_transcribe.assert_called_once_with("/tmp/telegram-busy-voice.ogg") + adapter.send.assert_awaited_once_with( + "12345", + '🎙️ "interrupt me"', + metadata={}, + ) + + def test_telegram_audio_size_gate_rejects_oversized_media_before_download(): adapter = object.__new__(TelegramAdapter) adapter._max_doc_bytes = 1024