[inworld] add falbback for empty timestamps from server

This commit is contained in:
Ian Lee
2026-03-20 17:18:54 -07:00
parent 4ee4002d5d
commit 5645909d34
2 changed files with 66 additions and 0 deletions

View File

@@ -0,0 +1 @@
- Fixed `InworldTtsService` to fallback to full text when TTS timestamps are not received

View File

@@ -28,6 +28,7 @@ from typing import (
Mapping, Mapping,
Optional, Optional,
Self, Self,
Set,
Tuple, Tuple,
) )
@@ -51,6 +52,7 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
from pipecat.frames.frames import ( from pipecat.frames.frames import (
AggregationType,
CancelFrame, CancelFrame,
EndFrame, EndFrame,
ErrorFrame, ErrorFrame,
@@ -60,6 +62,7 @@ from pipecat.frames.frames import (
TTSAudioRawFrame, TTSAudioRawFrame,
TTSStartedFrame, TTSStartedFrame,
TTSStoppedFrame, TTSStoppedFrame,
TTSTextFrame,
) )
from pipecat.processors.frame_processor import FrameDirection from pipecat.processors.frame_processor import FrameDirection
from pipecat.services.tts_service import TextAggregationMode, TTSService, WebsocketTTSService from pipecat.services.tts_service import TextAggregationMode, TTSService, WebsocketTTSService
@@ -215,6 +218,7 @@ class InworldHttpTTSService(TTSService):
self._base_url = "https://api.inworld.ai/tts/v1/voice" self._base_url = "https://api.inworld.ai/tts/v1/voice"
self._cumulative_time = 0.0 self._cumulative_time = 0.0
self._current_run_had_timestamps = False
# Init-only config (not runtime-updatable). # Init-only config (not runtime-updatable).
self._audio_encoding = encoding self._audio_encoding = encoding
@@ -299,6 +303,8 @@ class InworldHttpTTSService(TTSService):
""" """
logger.debug(f"{self}: Generating TTS [{text}] (streaming={self._streaming})") logger.debug(f"{self}: Generating TTS [{text}] (streaming={self._streaming})")
self._current_run_had_timestamps = False
audio_config = { audio_config = {
"audioEncoding": self._audio_encoding, "audioEncoding": self._audio_encoding,
"sampleRateHertz": self._audio_sample_rate, "sampleRateHertz": self._audio_sample_rate,
@@ -348,6 +354,23 @@ class InworldHttpTTSService(TTSService):
await self.start_tts_usage_metrics(text) await self.start_tts_usage_metrics(text)
# If no timestamps were received, push the full text so the LLM
# conversation context still reflects what the agent spoke. On
# interruption this means the full text is committed rather than
# only the portion that was spoken.
if not self._current_run_had_timestamps:
text_clean = text.rstrip()
if text_clean:
logger.debug(
f"{self}: No timestamps received, pushing fallback text: [{text_clean}]"
)
fallback = TTSTextFrame(
text_clean, aggregated_by=AggregationType.SENTENCE, context_id=context_id
)
ctx = self._tts_contexts.get(context_id)
fallback.append_to_context = ctx.append_to_context if ctx else True
await self.push_frame(fallback)
except Exception as e: except Exception as e:
await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e) await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e)
@@ -394,6 +417,7 @@ class InworldHttpTTSService(TTSService):
timestamp_info = chunk_data["result"]["timestampInfo"] timestamp_info = chunk_data["result"]["timestampInfo"]
word_times, chunk_end_time = self._calculate_word_times(timestamp_info) word_times, chunk_end_time = self._calculate_word_times(timestamp_info)
if word_times: if word_times:
self._current_run_had_timestamps = True
await self.add_word_timestamps(word_times, context_id) await self.add_word_timestamps(word_times, context_id)
# Track the maximum end time across all chunks # Track the maximum end time across all chunks
utterance_duration = max(utterance_duration, chunk_end_time) utterance_duration = max(utterance_duration, chunk_end_time)
@@ -430,6 +454,7 @@ class InworldHttpTTSService(TTSService):
timestamp_info = response_data["timestampInfo"] timestamp_info = response_data["timestampInfo"]
word_times, chunk_end_time = self._calculate_word_times(timestamp_info) word_times, chunk_end_time = self._calculate_word_times(timestamp_info)
if word_times: if word_times:
self._current_run_had_timestamps = True
await self.add_word_timestamps(word_times, context_id) await self.add_word_timestamps(word_times, context_id)
utterance_duration = chunk_end_time utterance_duration = chunk_end_time
@@ -658,6 +683,12 @@ class InworldTTSService(WebsocketTTSService):
# open contexts without causing duplicate creates in run_tts. # open contexts without causing duplicate creates in run_tts.
self._sent_context_ids: set[str] = set() self._sent_context_ids: set[str] = set()
# Fallback tracking for when timestamps are not received. Without
# timestamps, interruptions commit the full text rather than only the
# portion that was spoken.
self._context_texts: Dict[str, str] = {}
self._contexts_with_timestamps: Set[str] = set()
# Init-only config (not runtime-updatable). # Init-only config (not runtime-updatable).
self._audio_encoding = encoding self._audio_encoding = encoding
self._audio_sample_rate = 0 # Set in start() self._audio_sample_rate = 0 # Set in start()
@@ -803,8 +834,35 @@ class InworldTTSService(WebsocketTTSService):
async def on_audio_context_interrupted(self, context_id: str): async def on_audio_context_interrupted(self, context_id: str):
"""Callback invoked when an audio context has been interrupted.""" """Callback invoked when an audio context has been interrupted."""
await self._maybe_push_fallback_text(context_id)
await self._close_context(context_id) await self._close_context(context_id)
async def on_audio_context_completed(self, context_id: str):
"""Callback invoked when an audio context has been completed."""
await self._close_context(context_id)
async def _maybe_push_fallback_text(self, context_id: str):
"""Push the full text as fallback when no timestamps were received.
so that the LLM conversation context still reflects what the agent spoke.
Without timestamps, the full text is always committed — even on
interruption — since there is no timing information to determine which
portion was actually spoken.
"""
if not context_id:
return
had_timestamps = context_id in self._contexts_with_timestamps
text = self._context_texts.pop(context_id, "").strip()
self._contexts_with_timestamps.discard(context_id)
if had_timestamps or not text:
return
logger.debug(f"{self}: No timestamps for context {context_id}, pushing fallback: [{text}]")
fallback = TTSTextFrame(text, aggregated_by=AggregationType.SENTENCE)
fallback.context_id = context_id
ctx = self._tts_contexts.get(context_id)
fallback.append_to_context = ctx.append_to_context if ctx else True
await self.push_frame(fallback)
def _get_websocket(self): def _get_websocket(self):
"""Get the websocket for the Inworld WebSocket TTS service. """Get the websocket for the Inworld WebSocket TTS service.
@@ -913,6 +971,8 @@ class InworldTTSService(WebsocketTTSService):
self._sent_context_ids.clear() self._sent_context_ids.clear()
self._cumulative_time = 0.0 self._cumulative_time = 0.0
self._generation_end_time = 0.0 self._generation_end_time = 0.0
self._context_texts.clear()
self._contexts_with_timestamps.clear()
await self._call_event_handler("on_disconnected") await self._call_event_handler("on_disconnected")
async def _receive_messages(self): async def _receive_messages(self):
@@ -982,6 +1042,8 @@ class InworldTTSService(WebsocketTTSService):
if timestamp_info: if timestamp_info:
word_times = self._calculate_word_times(timestamp_info) word_times = self._calculate_word_times(timestamp_info)
if word_times: if word_times:
if ctx_id:
self._contexts_with_timestamps.add(ctx_id)
await self.add_word_timestamps(word_times, ctx_id) await self.add_word_timestamps(word_times, ctx_id)
# Handle flush completion, which indicates the end of a generation # Handle flush completion, which indicates the end of a generation
@@ -995,6 +1057,7 @@ class InworldTTSService(WebsocketTTSService):
# Handle context closed - context no longer exists on server # Handle context closed - context no longer exists on server
if "contextClosed" in result: if "contextClosed" in result:
logger.trace(f"{self}: Context closed on server: {ctx_id}") logger.trace(f"{self}: Context closed on server: {ctx_id}")
await self._maybe_push_fallback_text(ctx_id)
await self.stop_ttfb_metrics() await self.stop_ttfb_metrics()
await self.add_word_timestamps([("TTSStoppedFrame", 0), ("Reset", 0)], ctx_id) await self.add_word_timestamps([("TTSStoppedFrame", 0), ("Reset", 0)], ctx_id)
await self.remove_audio_context(ctx_id) await self.remove_audio_context(ctx_id)
@@ -1119,6 +1182,8 @@ class InworldTTSService(WebsocketTTSService):
yield TTSStartedFrame(context_id=context_id) yield TTSStartedFrame(context_id=context_id)
await self._send_context(context_id) await self._send_context(context_id)
self._context_texts[context_id] = self._context_texts.get(context_id, "") + text
await self._send_text(context_id, text) await self._send_text(context_id, text)
await self.start_tts_usage_metrics(text) await self.start_tts_usage_metrics(text)