Revert "feat(voice): start early replies before final asr"

This reverts commit a44a1b97a1.
This commit is contained in:
Yera All
2026-04-19 01:30:26 +05:00
parent 224973d840
commit aa64b6b94c
2 changed files with 100 additions and 469 deletions
+100 -333
View File
@@ -510,20 +510,6 @@ class AudioSocketMediaRuntime:
time.monotonic() + self._streaming_asr_reopen_backoff_seconds 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: async def _close_streaming_asr(self, actor: MediaActor, *, drain: bool = True) -> None:
stream_id = actor.asr_stream_id stream_id = actor.asr_stream_id
actor.asr_stream_id = None actor.asr_stream_id = None
@@ -565,10 +551,9 @@ class AudioSocketMediaRuntime:
bytes(batch), bytes(batch),
) )
except StreamingASRUnavailable as exc: except StreamingASRUnavailable as exc:
if actor.asr_stream_id == stream_id: actor.asr_streaming_failed = True
actor.asr_streaming_failed = True actor.asr_streaming_enabled = False
actor.asr_streaming_enabled = False self._mark_streaming_asr_backoff(actor)
self._mark_streaming_asr_backoff(actor)
logger.warning( logger.warning(
"audiosocket.streaming_asr_push_failed session_id=%s error=%s", "audiosocket.streaming_asr_push_failed session_id=%s error=%s",
actor.registration.voice_session_id, actor.registration.voice_session_id,
@@ -621,48 +606,6 @@ class AudioSocketMediaRuntime:
with contextlib.suppress(asyncio.CancelledError, Exception): with contextlib.suppress(asyncio.CancelledError, Exception):
await task 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: def _queue_streaming_asr_pcm(self, actor: MediaActor, pcm_frame: bytes) -> None:
queue = actor.streaming_asr_push_queue queue = actor.streaming_asr_push_queue
if queue is None or actor.asr_streaming_failed: 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) await self._set_actor_state(actor, "handoff_requested", final_decision.handoff_reason)
self._start_handoff_request(actor, transcript_text, final_decision) 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: async def _poll_streaming_partial(self, actor: MediaActor) -> None:
if actor.closed or not actor.asr_streaming_enabled or not actor.asr_stream_id: if actor.closed or not actor.asr_streaming_enabled or not actor.asr_stream_id:
return return
@@ -1033,80 +870,6 @@ class AudioSocketMediaRuntime:
finally: finally:
await self._close_streaming_asr(actor) 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( async def _record_reply_status(
self, self,
actor: MediaActor, actor: MediaActor,
@@ -1747,20 +1510,93 @@ class AudioSocketMediaRuntime:
ack_kind="unknown", ack_kind="unknown",
) )
use_v2 = self._should_use_voice_v2(actor.registration) transcript_source = "batch"
detached_stream = self._detach_streaming_asr(actor) if actor.asr_streaming_enabled else None if actor.asr_streaming_enabled and not actor.asr_streaming_failed:
final_transcription_task = asyncio.create_task( try:
self._finalize_turn_transcription( transcription = await self._finalize_streaming_transcription(actor)
actor, transcript_source = "streaming_final"
pcm_bytes=pcm_bytes, except StreamingASRUnavailable as exc:
partial_transcript=partial_transcript, logger.warning(
detached_stream=detached_stream, "audiosocket.streaming_asr_finalize_failed session_id=%s error=%s",
), actor.registration.voice_session_id,
name=f"voice-final-asr-{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 early_plan_decision: VoiceAITurnDecisionOut | None = None
if use_v2: if self._should_use_voice_v2(actor.registration):
early_plan_decision = await self._take_early_plan_decision(actor, timeout_seconds=0.12) early_plan_decision = await self._take_early_plan_decision(actor, timeout_seconds=0.04)
if early_plan_decision is not None: if early_plan_decision is not None:
decision = early_plan_decision decision = early_plan_decision
logger.info( logger.info(
@@ -1769,88 +1605,19 @@ class AudioSocketMediaRuntime:
actor.early_plan_generation, actor.early_plan_generation,
actor.early_plan_intent, actor.early_plan_intent,
) )
early_transcript_text = str(actor.early_plan_transcript or partial_transcript or "").strip() else:
early_metadata = { try:
**base_metadata, decision = await asyncio.wait_for(asyncio.shield(decision_task), timeout=self._v2_ack_wait_seconds)
**(decision.metadata if isinstance(decision.metadata, dict) else {}), except asyncio.TimeoutError:
"partial_transcript": early_transcript_text, if not actor.early_ack_started:
"early_intent": actor.early_plan_intent or partial_intent, await self._emit_early_ack(
"reply_phase": "early_plan", actor,
"transcript_source": "early_plan", language=transcription.language or actor.registration.language,
"early_reply_started_before_final_asr": True, metadata=metadata,
} ack_source="decision_timeout",
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,
) )
if remaining_gap > 0: early_plan_decision = await self._take_early_plan_decision(actor, timeout_seconds=0.12)
await asyncio.sleep(remaining_gap) decision = early_plan_decision if early_plan_decision is not None else await decision_task
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
else: else:
decision = await decision_task decision = await decision_task
handoff_task: asyncio.Task | None = None handoff_task: asyncio.Task | None = None
-136
View File
@@ -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") 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(): def test_media_runtime_voice_v2_uses_partial_as_final_when_streaming_finalize_fails():
captured: list[tuple[str, dict | None]] = [] captured: list[tuple[str, dict | None]] = []