From aa64b6b94c604dcd7eac898e7a211664e548f357 Mon Sep 17 00:00:00 2001 From: Yera All Date: Sun, 19 Apr 2026 01:30:26 +0500 Subject: [PATCH] Revert "feat(voice): start early replies before final asr" This reverts commit a44a1b97a105a63bbf9ad78d4b99fc8b9b313537. --- .../ai_voice_runtime_service/media_runtime.py | 433 ++++-------------- tests/test_ai_voice_media_runtime.py | 136 ------ 2 files changed, 100 insertions(+), 469 deletions(-) diff --git a/services/ai_voice_runtime_service/media_runtime.py b/services/ai_voice_runtime_service/media_runtime.py index 82ee255..fd88dd5 100644 --- a/services/ai_voice_runtime_service/media_runtime.py +++ b/services/ai_voice_runtime_service/media_runtime.py @@ -510,20 +510,6 @@ class AudioSocketMediaRuntime: time.monotonic() + self._streaming_asr_reopen_backoff_seconds ) - def _detach_streaming_asr( - self, - actor: MediaActor, - ) -> tuple[str | None, asyncio.Task | None, asyncio.Queue[bytes | None] | None, bool]: - stream_id = actor.asr_stream_id - push_task = actor.streaming_asr_push_task - push_queue = actor.streaming_asr_push_queue - streaming_failed = actor.asr_streaming_failed - actor.asr_stream_id = None - actor.asr_streaming_enabled = False - actor.streaming_asr_push_task = None - actor.streaming_asr_push_queue = None - return stream_id, push_task, push_queue, streaming_failed - async def _close_streaming_asr(self, actor: MediaActor, *, drain: bool = True) -> None: stream_id = actor.asr_stream_id actor.asr_stream_id = None @@ -565,10 +551,9 @@ class AudioSocketMediaRuntime: bytes(batch), ) except StreamingASRUnavailable as exc: - if actor.asr_stream_id == stream_id: - actor.asr_streaming_failed = True - actor.asr_streaming_enabled = False - self._mark_streaming_asr_backoff(actor) + actor.asr_streaming_failed = True + actor.asr_streaming_enabled = False + self._mark_streaming_asr_backoff(actor) logger.warning( "audiosocket.streaming_asr_push_failed session_id=%s error=%s", actor.registration.voice_session_id, @@ -621,48 +606,6 @@ class AudioSocketMediaRuntime: with contextlib.suppress(asyncio.CancelledError, Exception): await task - async def _stop_detached_streaming_asr_push_loop( - self, - actor: MediaActor, - *, - stream_id: str, - task: asyncio.Task | None, - queue: asyncio.Queue[bytes | None] | None, - drain: bool, - ) -> None: - if task is None: - return - if task.done(): - with contextlib.suppress(asyncio.CancelledError, Exception): - await task - return - if queue is not None: - if drain: - try: - await asyncio.wait_for( - queue.join(), - timeout=self._streaming_asr_push_drain_timeout_seconds, - ) - except asyncio.TimeoutError: - logger.warning( - "audiosocket.streaming_asr_push_drain_timeout session_id=%s stream_id=%s queued_frames=%s", - actor.registration.voice_session_id, - stream_id, - queue.qsize(), - ) - sentinel_enqueued = False - with contextlib.suppress(asyncio.QueueFull): - queue.put_nowait(None) - sentinel_enqueued = True - if not sentinel_enqueued: - task.cancel() - try: - await asyncio.wait_for(task, timeout=0.5) - except asyncio.TimeoutError: - task.cancel() - with contextlib.suppress(asyncio.CancelledError, Exception): - await task - def _queue_streaming_asr_pcm(self, actor: MediaActor, pcm_frame: bytes) -> None: queue = actor.streaming_asr_push_queue if queue is None or actor.asr_streaming_failed: @@ -857,112 +800,6 @@ class AudioSocketMediaRuntime: await self._set_actor_state(actor, "handoff_requested", final_decision.handoff_reason) self._start_handoff_request(actor, transcript_text, final_decision) - async def _start_final_decision_task( - self, - actor: MediaActor, - *, - transcription: ASRTranscription, - transcript_source: str, - partial_transcript: str, - base_metadata: dict[str, Any], - utterance_generation: int, - barge_in: bool, - set_listening_on_empty: bool, - ) -> tuple[asyncio.Task | None, str, dict[str, Any] | None]: - transcript_text = str(transcription.text or "").strip() or str(partial_transcript or "").strip() - logger.info( - "audiosocket.asr_turn_ready session_id=%s provider=%s utterance_ms=%s text_len=%s empty=%s", - actor.registration.voice_session_id, - getattr(self._asr_provider, "name", "unknown"), - int(base_metadata.get("turn_duration_ms") or 0), - len(str(transcript_text or "").strip()), - not bool(str(transcript_text or "").strip()), - ) - if not transcript_text: - if set_listening_on_empty: - await self._set_actor_state(actor, "listening") - return None, "", None - if self._is_low_signal_transcript(transcript_text): - actor.finalized_utterance_generation = max(actor.finalized_utterance_generation, utterance_generation) - logger.info( - "audiosocket.low_signal_ignored session_id=%s transcript=%s", - actor.registration.voice_session_id, - transcript_text[:120], - ) - if set_listening_on_empty: - await self._set_actor_state(actor, "listening") - return None, transcript_text, None - - actor.finalized_utterance_generation = max(actor.finalized_utterance_generation, utterance_generation) - actor.finalized_caller_turn_count += 1 - if actor.utterance_generation == utterance_generation: - actor.partial_transcript = transcript_text if actor.registration.voice_v2_partial_asr else None - actor.partial_intent = self._detect_early_intent(transcript_text) - actor.stable_partial_intent = actor.partial_intent - final_intent = self._detect_early_intent(transcript_text) - metadata = { - **base_metadata, - "partial_transcript": transcript_text if actor.registration.voice_v2_partial_asr else None, - "early_intent": final_intent, - "reply_phase": "final", - "transcript_source": transcript_source, - } - decision_task = asyncio.create_task( - asyncio.to_thread( - self._process_turn, - actor.registration.voice_session_id, - transcript_text, - transcription.language or actor.registration.language, - barge_in, - { - **metadata, - "runtime_defer_reply_planned": self._should_use_voice_v2(actor.registration), - }, - ), - name=f"voice-final-plan-{actor.registration.voice_session_id}", - ) - return decision_task, transcript_text, metadata - - async def _finalize_and_process_turn_after_early_reply( - self, - actor: MediaActor, - *, - final_transcription_task: asyncio.Task, - early_decision: VoiceAITurnDecisionOut, - partial_transcript: str, - base_metadata: dict[str, Any], - utterance_generation: int, - barge_in: bool, - ) -> None: - try: - transcription, transcript_source = await final_transcription_task - decision_task, transcript_text, _metadata = await self._start_final_decision_task( - actor, - transcription=transcription, - transcript_source=transcript_source, - partial_transcript=partial_transcript, - base_metadata=base_metadata, - utterance_generation=utterance_generation, - barge_in=barge_in, - set_listening_on_empty=False, - ) - if decision_task is None: - return - await self._reconcile_final_decision_after_early_plan( - actor, - decision_task=decision_task, - early_decision=early_decision, - transcript_text=transcript_text, - ) - except asyncio.CancelledError: - raise - except Exception as exc: - logger.warning( - "audiosocket.final_turn_background_failed session_id=%s error=%s", - actor.registration.voice_session_id, - str(exc)[:500], - ) - async def _poll_streaming_partial(self, actor: MediaActor) -> None: if actor.closed or not actor.asr_streaming_enabled or not actor.asr_stream_id: return @@ -1033,80 +870,6 @@ class AudioSocketMediaRuntime: finally: await self._close_streaming_asr(actor) - async def _finalize_detached_streaming_transcription( - self, - actor: MediaActor, - *, - stream_id: str, - push_task: asyncio.Task | None, - push_queue: asyncio.Queue[bytes | None] | None, - streaming_failed: bool, - ) -> ASRTranscription: - await self._stop_detached_streaming_asr_push_loop( - actor, - stream_id=stream_id, - task=push_task, - queue=push_queue, - drain=True, - ) - if streaming_failed: - raise StreamingASRUnavailable("Streaming ASR push failed before finalize") - try: - return await asyncio.to_thread(self._streaming_asr_provider.finalize, stream_id) - finally: - await asyncio.to_thread(self._streaming_asr_provider.close_stream, stream_id) - - async def _finalize_turn_transcription( - self, - actor: MediaActor, - *, - pcm_bytes: bytes, - partial_transcript: str, - detached_stream: tuple[str | None, asyncio.Task | None, asyncio.Queue[bytes | None] | None, bool] | None, - ) -> tuple[ASRTranscription, str]: - if detached_stream is not None and detached_stream[0]: - stream_id, push_task, push_queue, streaming_failed = detached_stream - try: - transcription = await self._finalize_detached_streaming_transcription( - actor, - stream_id=stream_id, - push_task=push_task, - push_queue=push_queue, - streaming_failed=streaming_failed, - ) - return transcription, "streaming_final" - except StreamingASRUnavailable as exc: - logger.warning( - "audiosocket.streaming_asr_finalize_failed session_id=%s error=%s", - actor.registration.voice_session_id, - str(exc)[:500], - ) - self._mark_streaming_asr_backoff(actor) - partial_first_text = str(partial_transcript or actor.stable_partial_transcript or actor.partial_transcript or "").strip() - if self._partial_first_final_enabled and partial_first_text and not self._is_low_signal_transcript(partial_first_text): - return ( - ASRTranscription( - text=partial_first_text, - language=actor.registration.language, - confidence=None, - ), - "streaming_partial_after_finalize_failure", - ) - wav_bytes = pcm16le_to_wav_bytes(pcm_bytes, sample_rate_hz=8000) - transcription = await asyncio.to_thread( - self._asr_provider.transcribe, - wav_bytes, - language_hint=actor.registration.language, - ) - return transcription, "batch_fallback_after_streaming_failure" - wav_bytes = pcm16le_to_wav_bytes(pcm_bytes, sample_rate_hz=8000) - transcription = await asyncio.to_thread( - self._asr_provider.transcribe, - wav_bytes, - language_hint=actor.registration.language, - ) - return transcription, "batch" - async def _record_reply_status( self, actor: MediaActor, @@ -1747,20 +1510,93 @@ class AudioSocketMediaRuntime: ack_kind="unknown", ) - use_v2 = self._should_use_voice_v2(actor.registration) - detached_stream = self._detach_streaming_asr(actor) if actor.asr_streaming_enabled else None - final_transcription_task = asyncio.create_task( - self._finalize_turn_transcription( - actor, - pcm_bytes=pcm_bytes, - partial_transcript=partial_transcript, - detached_stream=detached_stream, - ), - name=f"voice-final-asr-{actor.registration.voice_session_id}", + transcript_source = "batch" + if actor.asr_streaming_enabled and not actor.asr_streaming_failed: + try: + transcription = await self._finalize_streaming_transcription(actor) + transcript_source = "streaming_final" + except StreamingASRUnavailable as exc: + logger.warning( + "audiosocket.streaming_asr_finalize_failed session_id=%s error=%s", + actor.registration.voice_session_id, + str(exc)[:500], + ) + self._mark_streaming_asr_backoff(actor) + await self._close_streaming_asr(actor, drain=False) + partial_first_text = str(actor.stable_partial_transcript or actor.partial_transcript or "").strip() + if self._partial_first_final_enabled and partial_first_text and not self._is_low_signal_transcript(partial_first_text): + transcription = ASRTranscription( + text=partial_first_text, + language=actor.registration.language, + confidence=None, + ) + transcript_source = "streaming_partial_after_finalize_failure" + else: + wav_bytes = pcm16le_to_wav_bytes(pcm_bytes, sample_rate_hz=8000) + transcription = await asyncio.to_thread( + self._asr_provider.transcribe, + wav_bytes, + language_hint=actor.registration.language, + ) + transcript_source = "batch_fallback_after_streaming_failure" + else: + wav_bytes = pcm16le_to_wav_bytes(pcm_bytes, sample_rate_hz=8000) + transcription = await asyncio.to_thread( + self._asr_provider.transcribe, + wav_bytes, + language_hint=actor.registration.language, + ) + transcript_source = "batch" + transcript_text = str(transcription.text or "").strip() or partial_transcript + logger.info( + "audiosocket.asr_turn_ready session_id=%s provider=%s utterance_ms=%s text_len=%s empty=%s", + actor.registration.voice_session_id, + getattr(self._asr_provider, "name", "unknown"), + int(len(pcm_bytes) / 16), + len(str(transcript_text or "").strip()), + not bool(str(transcript_text or "").strip()), + ) + if not transcript_text: + await self._set_actor_state(actor, "listening") + return + if self._is_low_signal_transcript(transcript_text): + actor.finalized_utterance_generation = max(actor.finalized_utterance_generation, utterance_generation) + logger.info( + "audiosocket.low_signal_ignored session_id=%s transcript=%s", + actor.registration.voice_session_id, + transcript_text[:120], + ) + await self._set_actor_state(actor, "listening") + return + + actor.finalized_utterance_generation = max(actor.finalized_utterance_generation, utterance_generation) + actor.finalized_caller_turn_count += 1 + actor.partial_transcript = transcript_text if actor.registration.voice_v2_partial_asr else None + actor.partial_intent = self._detect_early_intent(transcript_text) + actor.stable_partial_intent = actor.partial_intent + metadata = { + **base_metadata, + "partial_transcript": actor.partial_transcript, + "early_intent": actor.partial_intent, + "reply_phase": "final", + "transcript_source": transcript_source, + } + decision_task = asyncio.create_task( + asyncio.to_thread( + self._process_turn, + actor.registration.voice_session_id, + transcript_text, + transcription.language or actor.registration.language, + barge_in, + { + **metadata, + "runtime_defer_reply_planned": self._should_use_voice_v2(actor.registration), + }, + ) ) early_plan_decision: VoiceAITurnDecisionOut | None = None - if use_v2: - early_plan_decision = await self._take_early_plan_decision(actor, timeout_seconds=0.12) + if self._should_use_voice_v2(actor.registration): + early_plan_decision = await self._take_early_plan_decision(actor, timeout_seconds=0.04) if early_plan_decision is not None: decision = early_plan_decision logger.info( @@ -1769,88 +1605,19 @@ class AudioSocketMediaRuntime: actor.early_plan_generation, actor.early_plan_intent, ) - early_transcript_text = str(actor.early_plan_transcript or partial_transcript or "").strip() - early_metadata = { - **base_metadata, - **(decision.metadata if isinstance(decision.metadata, dict) else {}), - "partial_transcript": early_transcript_text, - "early_intent": actor.early_plan_intent or partial_intent, - "reply_phase": "early_plan", - "transcript_source": "early_plan", - "early_reply_started_before_final_asr": True, - } - asyncio.create_task( - self._finalize_and_process_turn_after_early_reply( - actor, - final_transcription_task=final_transcription_task, - early_decision=early_plan_decision, - partial_transcript=early_transcript_text, - base_metadata=base_metadata, - utterance_generation=utterance_generation, - barge_in=barge_in, - ), - name=f"voice-final-turn-background-{actor.registration.voice_session_id}", - ) - handoff_task: asyncio.Task | None = None - if decision.needs_handoff: - await self._set_actor_state(actor, "handoff_requested", decision.handoff_reason) - handoff_task = self._start_handoff_request(actor, early_transcript_text, decision) - if decision.reply_text: - if actor.early_ack_started: - remaining_gap = self._v2_ack_post_gap_seconds - max( - 0.0, - time.monotonic() - actor.last_ack_completed_monotonic, + else: + try: + decision = await asyncio.wait_for(asyncio.shield(decision_task), timeout=self._v2_ack_wait_seconds) + except asyncio.TimeoutError: + if not actor.early_ack_started: + await self._emit_early_ack( + actor, + language=transcription.language or actor.registration.language, + metadata=metadata, + ack_source="decision_timeout", ) - if remaining_gap > 0: - await asyncio.sleep(remaining_gap) - await self._plan_reply_segment( - actor, - decision.reply_text, - kind="reply", - metadata={ - **early_metadata, - "phase": "main", - "early_ack_started": actor.early_ack_started, - }, - ) - await self._speak_reply(actor, decision.reply_text, is_greeting=False, reply_phase="main") - if actor.closed: - return - if decision.needs_handoff: - if handoff_task is not None and handoff_task.done(): - with contextlib.suppress(asyncio.CancelledError, Exception): - await handoff_task - return - actor.input_active = False - await self._set_actor_state(actor, "listening") - return - - transcription, transcript_source = await final_transcription_task - decision_task, transcript_text, metadata = await self._start_final_decision_task( - actor, - transcription=transcription, - transcript_source=transcript_source, - partial_transcript=partial_transcript, - base_metadata=base_metadata, - utterance_generation=utterance_generation, - barge_in=barge_in, - set_listening_on_empty=True, - ) - if decision_task is None or metadata is None: - return - if use_v2: - try: - decision = await asyncio.wait_for(asyncio.shield(decision_task), timeout=self._v2_ack_wait_seconds) - except asyncio.TimeoutError: - if not actor.early_ack_started: - await self._emit_early_ack( - actor, - language=transcription.language or actor.registration.language, - metadata=metadata, - ack_source="decision_timeout", - ) - early_plan_decision = await self._take_early_plan_decision(actor, timeout_seconds=0.12) - decision = early_plan_decision if early_plan_decision is not None else await decision_task + early_plan_decision = await self._take_early_plan_decision(actor, timeout_seconds=0.12) + decision = early_plan_decision if early_plan_decision is not None else await decision_task else: decision = await decision_task handoff_task: asyncio.Task | None = None diff --git a/tests/test_ai_voice_media_runtime.py b/tests/test_ai_voice_media_runtime.py index fc7ac2c..09407bf 100644 --- a/tests/test_ai_voice_media_runtime.py +++ b/tests/test_ai_voice_media_runtime.py @@ -1981,142 +1981,6 @@ def test_media_runtime_voice_v2_uses_early_plan_before_slow_final_decision(): assert events.index(("main", "early reply")) < events.index("final_done") -def test_media_runtime_voice_v2_starts_early_main_before_slow_final_asr(): - events: list[str | tuple[str, str]] = [] - - class _SlowStreamingProvider(StreamingASRProvider): - name = "slow-streaming" - supports_streaming = True - - def finalize(self, stream_id: str) -> ASRTranscription: - assert stream_id == "stream-1" - events.append("finalize_start") - time.sleep(0.35) - events.append("finalize_done") - return ASRTranscription(text="work schedule", language="ru", confidence=0.9) - - def close_stream(self, stream_id: str) -> None: - assert stream_id == "stream-1" - - def _process_turn(session_id, transcript_text, language, barge_in, metadata): - del session_id, transcript_text, barge_in - if metadata and metadata.get("reply_phase") == "early_plan": - events.append("early_plan") - return VoiceAITurnDecisionOut( - language=language or "ru", - intent="schedule", - reply_text="early reply", - confidence=0.8, - needs_handoff=False, - handoff_reason=None, - case_action="keep_open", - kb_refs=[], - summary_text="early ready", - model="early", - latency_ms=1, - status="active", - ) - events.append("final_turn") - return VoiceAITurnDecisionOut( - language=language or "ru", - intent="schedule", - reply_text="final reply", - confidence=0.9, - needs_handoff=False, - handoff_reason=None, - case_action="keep_open", - kb_refs=[], - summary_text="final ready", - model="final", - latency_ms=1, - status="active", - ) - - runtime = AudioSocketMediaRuntime( - enabled=True, - host="127.0.0.1", - port=0, - frame_ms=20, - idle_timeout_seconds=2.0, - registration_wait_timeout_seconds=0.5, - min_speech_ms=40, - trailing_silence_ms=40, - max_turn_ms=2000, - asr_provider=_StubASRProvider(), - streaming_asr_provider=_SlowStreamingProvider(), - tts_provider=_StubTTSProvider(), - load_registration_by_media_uuid=lambda value: None, - mark_media_connected=lambda session_id, value: None, - mark_media_ended=lambda session_id, reason: None, - touch_media_frame=lambda session_id: None, - set_state=lambda session_id, state, handoff_reason, metadata: None, - get_pending_greeting=lambda session_id: None, - mark_reply_delivered=lambda session_id, text, is_greeting: None, - plan_reply=lambda session_id, text, metadata, kind: None, - process_turn=_process_turn, - request_handoff=lambda session_id, customer_request_text, decision: None, - handle_media_error=lambda session_id, message, metadata: None, - ) - - async def _fake_speak_reply( - actor: MediaActor, - text: str, - *, - is_greeting: bool, - style_hints: dict[str, object] | None = None, - reply_phase: str | None = "main", - ) -> None: - del actor, is_greeting, style_hints - events.append((str(reply_phase), text)) - - runtime._speak_reply = _fake_speak_reply # type: ignore[method-assign] - - async def _scenario() -> None: - actor = MediaActor( - registration=MediaRegistration( - voice_session_id="avs_media_runtime_early_before_final_asr", - call_id="call_media_runtime_early_before_final_asr", - interaction_id="int_media_runtime_early_before_final_asr", - ai_session_id="ais_media_runtime_early_before_final_asr", - language="ru", - media_uuid=str(uuid.uuid4()), - queue_code="voice_lab_ai", - queue_id="que_voice_lab_ai", - agent_profile="voice_support", - voice_v2_enabled=True, - voice_v2_ack_mode="immediate_short", - voice_v2_streaming_tts=True, - voice_v2_partial_asr=True, - voice_v2_duplex=True, - voice_v2_streaming_asr_backend="local_sidecar", - ), - reader=asyncio.StreamReader(), - writer=None, # type: ignore[arg-type] - vad=EnergyVAD(frame_ms=20, min_speech_ms=40, trailing_silence_ms=40, max_turn_ms=2000), - frame_ms=20, - frame_bytes=320, - ) - actor.asr_streaming_enabled = True - actor.asr_stream_id = "stream-1" - actor.partial_transcript = "work schedule" - actor.stable_partial_transcript = "work schedule" - actor.partial_intent = "schedule" - actor.stable_partial_intent = "schedule" - pcm_frame = (1000).to_bytes(2, "little", signed=True) * 160 - await runtime._process_utterance(actor, pcm_frame * 40, False) - for _ in range(80): - if "final_turn" in events: - break - await asyncio.sleep(0.01) - - asyncio.run(_scenario()) - - assert "finalize_start" in events - assert ("main", "early reply") in events - assert events.index(("main", "early reply")) < events.index("finalize_done") - assert ("main", "final reply") not in events - - def test_media_runtime_voice_v2_uses_partial_as_final_when_streaming_finalize_fails(): captured: list[tuple[str, dict | None]] = []