Merge pull request #675 from pipecat-ai/mb/playht-add-request-id
Add a request_id to each TTS sequence
This commit is contained in:
@@ -8,6 +8,7 @@ import asyncio
|
|||||||
import io
|
import io
|
||||||
import json
|
import json
|
||||||
import struct
|
import struct
|
||||||
|
import uuid
|
||||||
from typing import AsyncGenerator, Optional
|
from typing import AsyncGenerator, Optional
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
@@ -127,6 +128,7 @@ class PlayHTTTSService(TTSService):
|
|||||||
self._websocket_url = None
|
self._websocket_url = None
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
self._request_id = None
|
||||||
|
|
||||||
self._settings = {
|
self._settings = {
|
||||||
"sample_rate": sample_rate,
|
"sample_rate": sample_rate,
|
||||||
@@ -191,6 +193,7 @@ class PlayHTTTSService(TTSService):
|
|||||||
await self._receive_task
|
await self._receive_task
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
|
self._request_id = None
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} error closing websocket: {e}")
|
logger.error(f"{self} error closing websocket: {e}")
|
||||||
|
|
||||||
@@ -221,6 +224,7 @@ class PlayHTTTSService(TTSService):
|
|||||||
async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
|
async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
|
||||||
await super()._handle_interruption(frame, direction)
|
await super()._handle_interruption(frame, direction)
|
||||||
await self.stop_all_metrics()
|
await self.stop_all_metrics()
|
||||||
|
self._request_id = None
|
||||||
|
|
||||||
async def _receive_task_handler(self):
|
async def _receive_task_handler(self):
|
||||||
try:
|
try:
|
||||||
@@ -242,9 +246,10 @@ class PlayHTTTSService(TTSService):
|
|||||||
logger.debug(f"Received text message: {message}")
|
logger.debug(f"Received text message: {message}")
|
||||||
try:
|
try:
|
||||||
msg = json.loads(message)
|
msg = json.loads(message)
|
||||||
if "request_id" in msg:
|
if "request_id" in msg and msg["request_id"] == self._request_id:
|
||||||
await self.push_frame(TTSStoppedFrame())
|
await self.push_frame(TTSStoppedFrame())
|
||||||
header_received = False # Reset for the next audio stream
|
header_received = False # Reset for the next audio stream
|
||||||
|
self._request_id = None
|
||||||
elif "error" in msg:
|
elif "error" in msg:
|
||||||
logger.error(f"{self} error: {msg}")
|
logger.error(f"{self} error: {msg}")
|
||||||
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
|
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
|
||||||
@@ -263,8 +268,10 @@ class PlayHTTTSService(TTSService):
|
|||||||
if not self._websocket or self._websocket.closed:
|
if not self._websocket or self._websocket.closed:
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
await self.start_ttfb_metrics()
|
if not self._request_id:
|
||||||
yield TTSStartedFrame()
|
await self.start_ttfb_metrics()
|
||||||
|
yield TTSStartedFrame()
|
||||||
|
self._request_id = str(uuid.uuid4())
|
||||||
|
|
||||||
tts_command = {
|
tts_command = {
|
||||||
"text": text,
|
"text": text,
|
||||||
@@ -275,6 +282,7 @@ class PlayHTTTSService(TTSService):
|
|||||||
"language": self._settings["language"],
|
"language": self._settings["language"],
|
||||||
"speed": self._settings["speed"],
|
"speed": self._settings["speed"],
|
||||||
"seed": self._settings["seed"],
|
"seed": self._settings["seed"],
|
||||||
|
"request_id": self._request_id,
|
||||||
}
|
}
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
Reference in New Issue
Block a user