Merge pull request #3288 from pipecat-ai/mb/inworld-cleanup

Inworld TTS service clean up
This commit is contained in:
Mark Backman
2025-12-29 13:07:20 -05:00
committed by GitHub
2 changed files with 86 additions and 54 deletions

View File

@@ -0,0 +1,5 @@
- Updates to Inworld TTS services:
- Improved `InworldTTSService`'s websocket implementation to better flush and
close context to better handle long inputs.
- Improved docstrings for `InworldTTSService` and `InworldHttpTTSService`.

View File

@@ -7,8 +7,10 @@
"""Inworld AI Text-to-Speech Service Implementation. """Inworld AI Text-to-Speech Service Implementation.
Contains two TTS services: Contains two TTS services:
- InworldHttpTTSService: HTTP-based TTS service.
- InworldTTSService: WebSocket-based TTS service. - InworldTTSService: WebSocket-based TTS service.
- InworldHttpTTSService: HTTP-based TTS service.
Inworlds text-to-speech (TTS) models offer ultra-realistic, context-aware speech synthesis and precise voice cloning capabilities, enabling developers to build natural and engaging experiences with human-like speech quality at an accessible price point.
""" """
import base64 import base64
@@ -48,7 +50,7 @@ class InworldHttpTTSService(WordTTSService):
"""Inworld AI HTTP-based TTS service. """Inworld AI HTTP-based TTS service.
Supports both streaming and non-streaming modes via the `streaming` parameter. Supports both streaming and non-streaming modes via the `streaming` parameter.
Outputs LINEAR16 audio at configurable sample rates with word/character timestamps. Outputs LINEAR16 audio at configurable sample rates with word-level timestamps.
""" """
class InputParams(BaseModel): class InputParams(BaseModel):
@@ -279,7 +281,7 @@ class InworldHttpTTSService(WordTTSService):
Args: Args:
response: The response from the Inworld API. response: The response from the Inworld API.
Returns: Yields:
An asynchronous generator of frames. An asynchronous generator of frames.
""" """
buffer = "" buffer = ""
@@ -398,7 +400,7 @@ class InworldTTSService(AudioContextWordTTSService):
Uses bidirectional WebSocket for lower latency streaming. Supports multiple Uses bidirectional WebSocket for lower latency streaming. Supports multiple
independent audio contexts per connection (max 5). Outputs LINEAR16 audio independent audio contexts per connection (max 5). Outputs LINEAR16 audio
with word/character timestamps. with word-level timestamps.
""" """
class InputParams(BaseModel): class InputParams(BaseModel):
@@ -520,15 +522,15 @@ class InworldTTSService(AudioContextWordTTSService):
await self._disconnect() await self._disconnect()
async def flush_audio(self): async def flush_audio(self):
"""Flush any pending audio from the Inworld WebSocket TTS service. """Flush any pending audio without closing the context.
Args: This triggers synthesis of all accumulated text in the buffer while
frame: The flush frame. keeping the context open for subsequent text. The context is only
closed on interruption, disconnect, or end of session.
""" """
if self._context_id: if self._context_id and self._websocket:
ctx_to_close = self._context_id logger.trace(f"Flushing audio for context {self._context_id}")
self._context_id = None await self._send_flush(self._context_id)
await self._send_close_context(ctx_to_close)
async def push_frame(self, frame: Frame, direction: FrameDirection = FrameDirection.DOWNSTREAM): async def push_frame(self, frame: Frame, direction: FrameDirection = FrameDirection.DOWNSTREAM):
"""Push a frame and handle state changes. """Push a frame and handle state changes.
@@ -538,16 +540,14 @@ class InworldTTSService(AudioContextWordTTSService):
direction: The direction to push the frame. direction: The direction to push the frame.
""" """
await super().push_frame(frame, direction) await super().push_frame(frame, direction)
if isinstance(frame, TTSStoppedFrame): if isinstance(frame, (TTSStoppedFrame, InterruptionFrame)):
self._started = False self._started = False
await self.add_word_timestamps([("Reset", 0)]) if isinstance(frame, TTSStoppedFrame):
await self.add_word_timestamps([("Reset", 0)])
def _calculate_word_times(self, timestamp_info: Dict[str, Any]) -> List[Tuple[str, float]]: def _calculate_word_times(self, timestamp_info: Dict[str, Any]) -> List[Tuple[str, float]]:
"""Calculate word timestamps from Inworld WebSocket API response. """Calculate word timestamps from Inworld WebSocket API response.
Note: Inworld WebSocket provides cumulative timestamps across all chunks
within a conversation turn, similar to Cartesia. No additional tracking needed.
Args: Args:
timestamp_info: The timestamp information from Inworld API. timestamp_info: The timestamp information from Inworld API.
@@ -573,16 +573,21 @@ class InworldTTSService(AudioContextWordTTSService):
frame: The interruption frame. frame: The interruption frame.
direction: The direction of the interruption. direction: The direction of the interruption.
""" """
old_context_id = self._context_id
logger.trace(f"{self}: Handling interruption, old context: {old_context_id}")
await super()._handle_interruption(frame, direction) await super()._handle_interruption(frame, direction)
if self._context_id and self._websocket: if old_context_id and self._websocket:
logger.trace(f"Closing context {self._context_id} due to interruption") logger.trace(f"{self}: Closing context {old_context_id} due to interruption")
try: try:
await self._send_close_context(self._context_id) await self._send_close_context(old_context_id)
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)
self._context_id = None
self._started = False self._context_id = None
self._started = False
logger.trace(f"{self}: Interruption handled, context reset to None")
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.
@@ -658,16 +663,11 @@ class InworldTTSService(AudioContextWordTTSService):
finally: finally:
self._started = False self._started = False
self._context_id = None self._context_id = None
self._cumulative_time = 0.0
self._websocket = None self._websocket = None
await self._call_event_handler("on_disconnected") await self._call_event_handler("on_disconnected")
async def _process_messages(self): async def _receive_messages(self):
"""Process incoming WebSocket messages from Inworld. """Handle incoming WebSocket messages from Inworld."""
Returns:
The messages.
"""
async for message in self._get_websocket(): async for message in self._get_websocket():
try: try:
msg = json.loads(message) msg = json.loads(message)
@@ -678,6 +678,17 @@ class InworldTTSService(AudioContextWordTTSService):
result = msg.get("result", {}) result = msg.get("result", {})
ctx_id = result.get("contextId") or result.get("context_id") ctx_id = result.get("contextId") or result.get("context_id")
# Log all incoming messages for debugging
msg_types = [
k
for k in ["contextCreated", "audioChunk", "flushCompleted", "contextClosed"]
if k in result
]
logger.debug(
f"{self}: Received message types={msg_types}, ctx_id={ctx_id}, "
f"current_ctx={self._context_id}, available={self.audio_context_available(ctx_id) if ctx_id else 'N/A'}"
)
# Check for errors # Check for errors
status = result.get("status", {}) status = result.get("status", {})
if status.get("code", 0) != 0: if status.get("code", 0) != 0:
@@ -689,15 +700,26 @@ class InworldTTSService(AudioContextWordTTSService):
await self.push_error(error_msg=str(msg["error"])) await self.push_error(error_msg=str(msg["error"]))
continue continue
# Skip messages for unavailable contexts # Check if this message belongs to an available context.
# If the context isn't available but matches our current context ID,
# recreate it (handles race conditions during interruption recovery).
if ctx_id and not self.audio_context_available(ctx_id): if ctx_id and not self.audio_context_available(ctx_id):
continue if self._context_id == ctx_id:
logger.trace(
f"{self}: Recreating audio context for current context: {self._context_id}"
)
await self.create_audio_context(self._context_id)
else:
# This is a message from an old/closed context - skip it
logger.trace(f"{self}: Skipping message from unavailable context: {ctx_id}")
continue
# Process audio chunk # Process audio chunk
audio_chunk = result.get("audioChunk", {}) audio_chunk = result.get("audioChunk", {})
audio_b64 = audio_chunk.get("audioContent") audio_b64 = audio_chunk.get("audioContent")
if audio_b64: if audio_b64:
logger.trace(f"{self}: Processing audio chunk for context {ctx_id}")
await self.stop_ttfb_metrics() await self.stop_ttfb_metrics()
await self.start_word_timestamps() await self.start_word_timestamps()
audio = base64.b64decode(audio_b64) audio = base64.b64decode(audio_b64)
@@ -706,8 +728,6 @@ class InworldTTSService(AudioContextWordTTSService):
frame = TTSAudioRawFrame(audio, self.sample_rate, 1) frame = TTSAudioRawFrame(audio, self.sample_rate, 1)
if ctx_id: if ctx_id:
if not self.audio_context_available(ctx_id):
await self.create_audio_context(ctx_id)
await self.append_to_audio_context(ctx_id, frame) await self.append_to_audio_context(ctx_id, frame)
# timestampInfo is inside audioChunk # timestampInfo is inside audioChunk
@@ -717,24 +737,25 @@ class InworldTTSService(AudioContextWordTTSService):
if word_times: if word_times:
await self.add_word_timestamps(word_times) await self.add_word_timestamps(word_times)
# Handle context completion # Handle context created confirmation
if "flushCompleted" in result or "contextClosed" in result: if "contextCreated" in result:
logger.trace(f"{self}: Context created on server: {ctx_id}")
# Handle flush completion - context is still valid, just acknowledge it
if "flushCompleted" in result:
logger.trace(f"{self}: Flush completed for context {ctx_id}")
# Handle context closed - context no longer exists on server
if "contextClosed" in result:
logger.trace(f"{self}: Context closed on server: {ctx_id}")
await self.stop_ttfb_metrics() await self.stop_ttfb_metrics()
await self.add_word_timestamps([("TTSStoppedFrame", 0), ("Reset", 0)]) # Only reset if this is our current context
if ctx_id == self._context_id:
self._context_id = None
self._started = False
if ctx_id and self.audio_context_available(ctx_id): if ctx_id and self.audio_context_available(ctx_id):
await self.remove_audio_context(ctx_id) await self.remove_audio_context(ctx_id)
await self.add_word_timestamps([("TTSStoppedFrame", 0), ("Reset", 0)])
async def _receive_messages(self):
"""Receive messages from the Inworld WebSocket TTS service with auto-reconnect.
Returns:
The messages.
"""
while True:
await self._process_messages()
# Inworld may disconnect after period of inactivity, so we try to reconnect
logger.debug(f"{self} Inworld connection was disconnected, reconnecting")
await self._connect_websocket()
async def _send_context(self, context_id: str): async def _send_context(self, context_id: str):
"""Send a context to the Inworld WebSocket TTS service. """Send a context to the Inworld WebSocket TTS service.
@@ -752,14 +773,16 @@ class InworldTTSService(AudioContextWordTTSService):
create_config["temperature"] = self._settings["temperature"] create_config["temperature"] = self._settings["temperature"]
if "applyTextNormalization" in self._settings: if "applyTextNormalization" in self._settings:
create_config["applyTextNormalization"] = self._settings["applyTextNormalization"] create_config["applyTextNormalization"] = self._settings["applyTextNormalization"]
if self._buffer_settings["maxBufferDelayMs"] is not None:
create_config["maxBufferDelayMs"] = self._buffer_settings["maxBufferDelayMs"] # Set buffer settings for timely audio generation.
if self._buffer_settings["bufferCharThreshold"] is not None: # Use provided values or defaults that work well for streaming LLM output.
create_config["bufferCharThreshold"] = self._buffer_settings["bufferCharThreshold"] create_config["maxBufferDelayMs"] = self._buffer_settings["maxBufferDelayMs"] or 3000
create_config["bufferCharThreshold"] = self._buffer_settings["bufferCharThreshold"] or 250
create_config["timestampType"] = self._timestamp_type create_config["timestampType"] = self._timestamp_type
msg = {"create": create_config, "contextId": context_id} msg = {"create": create_config, "contextId": context_id}
logger.trace(f"{self}: Sending context create: {create_config}")
await self.send_with_retry(json.dumps(msg), self._report_error) await self.send_with_retry(json.dumps(msg), self._report_error)
async def _send_text(self, context_id: str, text: str): async def _send_text(self, context_id: str, text: str):
@@ -811,12 +834,16 @@ class InworldTTSService(AudioContextWordTTSService):
await self.start_ttfb_metrics() await self.start_ttfb_metrics()
yield TTSStartedFrame() yield TTSStartedFrame()
self._started = True self._started = True
if not self._context_id: if not self._context_id:
self._context_id = str(uuid.uuid4()) self._context_id = str(uuid.uuid4())
if not self.audio_context_available(self._context_id): logger.trace(f"{self}: Creating new context {self._context_id}")
await self.create_audio_context(self._context_id)
await self._send_context(self._context_id)
elif not self.audio_context_available(self._context_id):
# Context exists on server but local tracking was removed
logger.trace(f"{self}: Recreating local audio context {self._context_id}")
await self.create_audio_context(self._context_id) await self.create_audio_context(self._context_id)
await self._send_context(self._context_id)
await self._send_text(self._context_id, text) await self._send_text(self._context_id, text)
await self.start_tts_usage_metrics(text) await self.start_tts_usage_metrics(text)