add on_turn_context_created hook instead

This commit is contained in:
Ian Lee
2026-03-20 10:33:25 -07:00
parent dfe5fec8f9
commit e9f3086ea3
2 changed files with 35 additions and 35 deletions

View File

@@ -56,7 +56,6 @@ from pipecat.frames.frames import (
ErrorFrame, ErrorFrame,
Frame, Frame,
InterruptionFrame, InterruptionFrame,
LLMFullResponseStartFrame,
StartFrame, StartFrame,
TTSAudioRawFrame, TTSAudioRawFrame,
TTSStartedFrame, TTSStartedFrame,
@@ -654,10 +653,10 @@ class InworldTTSService(WebsocketTTSService):
# Track the end time of the last word in the current generation # Track the end time of the last word in the current generation
self._generation_end_time = 0.0 self._generation_end_time = 0.0
# Context ID that was pre-opened on the server during process_frame # Context IDs already sent to the server via _send_context, used to
# (LLMFullResponseStartFrame) to avoid context creation latency when # make _send_context idempotent so on_turn_context_created can eagerly
# enough context for TTS is available. # open contexts without causing duplicate creates in run_tts.
self._prewarmed_context_id: Optional[str] = None self._sent_context_ids: set[str] = set()
# Init-only config (not runtime-updatable). # Init-only config (not runtime-updatable).
self._audio_encoding = encoding self._audio_encoding = encoding
@@ -732,28 +731,16 @@ class InworldTTSService(WebsocketTTSService):
if isinstance(frame, TTSStoppedFrame): if isinstance(frame, TTSStoppedFrame):
await self.add_word_timestamps([("Reset", 0)]) await self.add_word_timestamps([("Reset", 0)])
async def process_frame(self, frame: Frame, direction: FrameDirection): async def on_turn_context_created(self, context_id: str):
"""Process incoming frames and pre-open context on LLM response start. """Eagerly open the context on the server when a new turn starts.
Eagerly sends the context configuration to the server when This overlaps server-side context creation with sentence aggregation
LLMFullResponseStartFrame arrives, so the context is ready by the time time, so the context is ready by the time text arrives in run_tts.
enough context for TTS is available. The base class assigns ``_turn_context_id`` before
this runs, which is reused for all ``run_tts`` calls within the turn.
""" """
await super().process_frame(frame, direction) try:
await self._send_context(context_id)
if isinstance(frame, LLMFullResponseStartFrame): except Exception as e:
if self._prewarmed_context_id: logger.warning(f"{self}: Failed to pre-open context: {e}")
try:
await self._send_close_context(self._prewarmed_context_id)
except Exception as e:
logger.warning(f"{self}: Failed to close previous prewarmed context: {e}")
self._prewarmed_context_id = None
try:
await self._send_context(self._turn_context_id)
self._prewarmed_context_id = self._turn_context_id
except Exception as e:
logger.warning(f"{self}: Failed to pre-open context: {e}")
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.
@@ -800,6 +787,7 @@ class InworldTTSService(WebsocketTTSService):
await self._send_close_context(context_id) await self._send_close_context(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._sent_context_ids.discard(context_id)
self._cumulative_time = 0.0 self._cumulative_time = 0.0
self._generation_end_time = 0.0 self._generation_end_time = 0.0
@@ -916,7 +904,7 @@ class InworldTTSService(WebsocketTTSService):
finally: finally:
await self.remove_active_audio_context() await self.remove_active_audio_context()
self._websocket = None self._websocket = None
self._prewarmed_context_id = None 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
await self._call_event_handler("on_disconnected") await self._call_event_handler("on_disconnected")
@@ -961,8 +949,11 @@ class InworldTTSService(WebsocketTTSService):
await self.push_error(error_msg=str(msg["error"])) await self.push_error(error_msg=str(msg["error"]))
continue continue
# Handle context created confirmation
if "contextCreated" in result:
logger.trace(f"{self}: Context created on server: {ctx_id}")
# If the context isn't available recreate it (handles race conditions during interruption recovery). # If the context isn't available recreate it (handles race conditions during interruption recovery).
if ctx_id and not self.audio_context_available(ctx_id): elif ctx_id and not self.audio_context_available(ctx_id):
logger.trace(f"{self}: Recreating audio context for current context: {ctx_id}") logger.trace(f"{self}: Recreating audio context for current context: {ctx_id}")
await self.create_audio_context(ctx_id) await self.create_audio_context(ctx_id)
@@ -987,10 +978,6 @@ class InworldTTSService(WebsocketTTSService):
if word_times: if word_times:
await self.add_word_timestamps(word_times, ctx_id) await self.add_word_timestamps(word_times, ctx_id)
# Handle context created confirmation
if "contextCreated" in result:
logger.trace(f"{self}: Context created on server: {ctx_id}")
# Handle flush completion, which indicates the end of a generation # Handle flush completion, which indicates the end of a generation
if "flushCompleted" in result: if "flushCompleted" in result:
logger.trace( logger.trace(
@@ -1031,15 +1018,15 @@ class InworldTTSService(WebsocketTTSService):
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.
Skips the send if this context was already pre-opened on the server Idempotent: skips the send if this context was already opened on the
(prewarmed during process_frame). server (e.g., eagerly via on_turn_context_created).
Args: Args:
context_id: The context ID. context_id: The context ID.
""" """
if context_id == self._prewarmed_context_id: if context_id in self._sent_context_ids:
self._prewarmed_context_id = None
return return
self._sent_context_ids.add(context_id)
audio_config = { audio_config = {
"audioEncoding": self._audio_encoding, "audioEncoding": self._audio_encoding,

View File

@@ -633,6 +633,18 @@ class TTSService(AIService):
await self.queue_frame(TTSSpeakFrame(text)) await self.queue_frame(TTSSpeakFrame(text))
async def on_turn_context_created(self, context_id: str):
"""Called when a new turn context ID has been created.
Override to perform provider-specific setup (e.g., eagerly opening a
server-side context) before text starts flowing. This is called from
``process_frame`` when an ``LLMFullResponseStartFrame`` arrives.
Args:
context_id: The newly created turn context ID.
"""
pass
async def on_turn_context_completed(self): async def on_turn_context_completed(self):
"""Handle the completion of a turn.""" """Handle the completion of a turn."""
# For HTTP services they emit the frames synchronously, so close the audio context here # For HTTP services they emit the frames synchronously, so close the audio context here
@@ -685,6 +697,7 @@ class TTSService(AIService):
self._llm_response_started = True self._llm_response_started = True
# New LLM turn → assign a fresh context ID shared by all sentences # New LLM turn → assign a fresh context ID shared by all sentences
self._turn_context_id = self.create_context_id() self._turn_context_id = self.create_context_id()
await self.on_turn_context_created(self._turn_context_id)
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
elif isinstance(frame, (LLMFullResponseEndFrame, EndFrame)): elif isinstance(frame, (LLMFullResponseEndFrame, EndFrame)):
# We pause processing incoming frames if the LLM response included # We pause processing incoming frames if the LLM response included