cartesia websockets and streaming
This commit is contained in:
@@ -149,6 +149,10 @@ class TTSService(AIService):
|
|||||||
async def say(self, text: str):
|
async def say(self, text: str):
|
||||||
await self.process_frame(TextFrame(text=text), FrameDirection.DOWNSTREAM)
|
await self.process_frame(TextFrame(text=text), FrameDirection.DOWNSTREAM)
|
||||||
|
|
||||||
|
async def handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
|
||||||
|
self._current_sentence = ""
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
async def _process_text_frame(self, frame: TextFrame):
|
async def _process_text_frame(self, frame: TextFrame):
|
||||||
text: str | None = None
|
text: str | None = None
|
||||||
if not self._aggregate_sentences:
|
if not self._aggregate_sentences:
|
||||||
@@ -182,8 +186,7 @@ class TTSService(AIService):
|
|||||||
if isinstance(frame, TextFrame):
|
if isinstance(frame, TextFrame):
|
||||||
await self._process_text_frame(frame)
|
await self._process_text_frame(frame)
|
||||||
elif isinstance(frame, StartInterruptionFrame):
|
elif isinstance(frame, StartInterruptionFrame):
|
||||||
self._current_sentence = ""
|
await self.handle_interruption(frame, direction)
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
elif isinstance(frame, LLMFullResponseEndFrame) or isinstance(frame, EndFrame):
|
elif isinstance(frame, LLMFullResponseEndFrame) or isinstance(frame, EndFrame):
|
||||||
self._current_sentence = ""
|
self._current_sentence = ""
|
||||||
await self._push_tts_frames(self._current_sentence)
|
await self._push_tts_frames(self._current_sentence)
|
||||||
|
|||||||
@@ -11,7 +11,8 @@ import asyncio
|
|||||||
|
|
||||||
from typing import AsyncGenerator
|
from typing import AsyncGenerator
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, CancelFrame, EndFrame, Frame, StartFrame
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
|
from pipecat.frames.frames import Frame, AudioRawFrame, StartInterruptionFrame
|
||||||
from pipecat.services.ai_services import TTSService
|
from pipecat.services.ai_services import TTSService
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -39,7 +40,7 @@ class CartesiaTTSService(TTSService):
|
|||||||
encoding: str = "pcm_s16le",
|
encoding: str = "pcm_s16le",
|
||||||
sample_rate: int = 16000,
|
sample_rate: int = 16000,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
super().__init__(aggregate_sentences=True, **kwargs)
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._cartesia_version = cartesia_version
|
self._cartesia_version = cartesia_version
|
||||||
@@ -53,31 +54,49 @@ class CartesiaTTSService(TTSService):
|
|||||||
}
|
}
|
||||||
self._language = "en"
|
self._language = "en"
|
||||||
|
|
||||||
|
self._websocket = None
|
||||||
self._context_id = None
|
self._context_id = None
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
self._waiting_for_ttfb = False
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def connect(self):
|
||||||
try:
|
try:
|
||||||
self._websocket = await websockets.connect(
|
self._websocket = await websockets.connect(
|
||||||
f"{self._url}?api_key={self._api_key}&cartesia_version={self._cartesia_version}"
|
f"{self._url}?api_key={self._api_key}&cartesia_version={self._cartesia_version}"
|
||||||
)
|
)
|
||||||
# self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"{self} initialization error: {e}")
|
logger.exception(f"{self} initialization error: {e}")
|
||||||
|
|
||||||
|
async def disconnect(self):
|
||||||
|
try:
|
||||||
|
if self._websocket:
|
||||||
|
ws = self._websocket
|
||||||
|
self._websocket = None
|
||||||
|
await ws.close()
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"{self} error closing websocket: {e}")
|
||||||
|
|
||||||
|
async def handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
|
||||||
|
await super().handle_interruption(frame, direction)
|
||||||
|
if self._receive_task:
|
||||||
|
self._receive_task.cancel()
|
||||||
|
self._receive_task = None
|
||||||
|
await self.disconnect()
|
||||||
|
await self.stop_all_metrics()
|
||||||
|
|
||||||
async def _receive_task_handler(self):
|
async def _receive_task_handler(self):
|
||||||
logger.debug("TOP OF RECEIVE TASK ...")
|
|
||||||
async for message in self._websocket:
|
async for message in self._websocket:
|
||||||
logger.debug("RECEIVE TASK LOOP")
|
|
||||||
msg = json.loads(message)
|
msg = json.loads(message)
|
||||||
if not msg:
|
if not msg:
|
||||||
continue
|
continue
|
||||||
logger.debug(f"Received message: {msg}")
|
# logger.debug(f"Received message: {msg}")
|
||||||
|
if self._waiting_for_ttfb:
|
||||||
|
await self.stop_ttfb_metrics()
|
||||||
|
self._waiting_for_ttfb = False
|
||||||
if msg["done"]:
|
if msg["done"]:
|
||||||
logger.debug(f"This was a 'done' message, shut down the receive task.")
|
|
||||||
self._context_id = None
|
self._context_id = None
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
self._receive_task.cancel()
|
self._receive_task.cancel()
|
||||||
@@ -90,96 +109,38 @@ class CartesiaTTSService(TTSService):
|
|||||||
)
|
)
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
# async for message in self._websocket:
|
|
||||||
# utterance = json.loads(message)
|
|
||||||
# if not utterance:
|
|
||||||
# continue
|
|
||||||
|
|
||||||
# logger.debug(f"Received utterance: {utterance}")
|
|
||||||
# return
|
|
||||||
|
|
||||||
# # TODO: PORT FROM GLADIA
|
|
||||||
# if "error" in utterance:
|
|
||||||
# message = utterance["message"]
|
|
||||||
# logger.error(f"Gladia error: {message}")
|
|
||||||
# elif "confidence" in utterance:
|
|
||||||
# type = utterance["type"]
|
|
||||||
# confidence = utterance["confidence"]
|
|
||||||
# transcript = utterance["transcription"]
|
|
||||||
# if confidence >= self._confidence:
|
|
||||||
# if type == "final":
|
|
||||||
# await self.queue_frame(TranscriptionFrame(transcript, "", int(time.time_ns() / 1000000)))
|
|
||||||
# else:
|
|
||||||
# await self.queue_frame(InterimTranscriptionFrame(transcript, "",
|
|
||||||
# int(time.time_ns() / 1000000)))
|
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
|
||||||
self._context_id = None
|
|
||||||
if self._receive_task:
|
|
||||||
self._receive_task.cancel()
|
|
||||||
self._receive_task = None
|
|
||||||
return
|
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
|
||||||
self._context_id = None
|
|
||||||
if self._receive_task:
|
|
||||||
self._receive_task.cancel()
|
|
||||||
self._receive_task = None
|
|
||||||
return
|
|
||||||
|
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
logger.debug(f"Generating TTS: [{text}]")
|
logger.debug(f"Generating TTS: [{text}]")
|
||||||
logger.debug(
|
|
||||||
f"model_id: {self._model_id}, voice_id: {self._voice_id}, language: {self._language}"
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await self.start_ttfb_metrics()
|
if not self._websocket:
|
||||||
|
await self.connect()
|
||||||
|
|
||||||
|
if not self._waiting_for_ttfb:
|
||||||
|
await self.start_ttfb_metrics()
|
||||||
|
self._waiting_for_ttfb = True
|
||||||
|
|
||||||
if not self._context_id:
|
if not self._context_id:
|
||||||
self._context_id = str(uuid.uuid4())
|
self._context_id = str(uuid.uuid4())
|
||||||
msg = {
|
|
||||||
"transcript": text,
|
msg = {
|
||||||
"continue": True,
|
"transcript": text,
|
||||||
"context_id": self._context_id,
|
"continue": True,
|
||||||
"model_id": self._model_id,
|
"context_id": self._context_id,
|
||||||
"voice": {
|
"model_id": self._model_id,
|
||||||
"mode": "id",
|
"voice": {
|
||||||
"id": self._voice_id
|
"mode": "id",
|
||||||
},
|
"id": self._voice_id
|
||||||
"output_format": self._output_format,
|
},
|
||||||
"language": self._language,
|
"output_format": self._output_format,
|
||||||
}
|
"language": self._language,
|
||||||
logger.debug(f"SENDING FIRST MESSAGE {json.dumps(msg)}")
|
}
|
||||||
await self._websocket.send(json.dumps(msg))
|
# logger.debug(f"SENDING MESSAGE {json.dumps(msg)}")
|
||||||
logger.debug("AWAITING FIRST RESPONSE MESSAGE")
|
await self._websocket.send(json.dumps(msg))
|
||||||
message = await self._websocket.recv()
|
if not self._receive_task:
|
||||||
msg = json.loads(message)
|
# todo: how do we await this task at the app level, so the program doesn't exit?
|
||||||
logger.debug(f"Received message: {msg}")
|
# we can't await here because we need this function to return
|
||||||
if (msg["type"] == "error"):
|
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
||||||
logger.error(f"Error: {msg}")
|
yield None
|
||||||
return
|
|
||||||
frame = AudioRawFrame(
|
|
||||||
audio=base64.b64decode(msg["data"]),
|
|
||||||
sample_rate=self._output_format["sample_rate"],
|
|
||||||
num_channels=1
|
|
||||||
)
|
|
||||||
yield frame
|
|
||||||
if not msg["done"]:
|
|
||||||
logger.debug("CREATING RECEIVE TASK")
|
|
||||||
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
|
||||||
# todo: how do we await this task at the app level, so the program doesn't exit?
|
|
||||||
# we can't await here because we need this function to return
|
|
||||||
# await self._receive_task
|
|
||||||
else:
|
|
||||||
msg = {
|
|
||||||
"transcript": text,
|
|
||||||
"continue": True,
|
|
||||||
"context_id": self._context_id,
|
|
||||||
}
|
|
||||||
await asyncio.sleep(0.350)
|
|
||||||
logger.debug(f"SENDING FOLLOW MESSAGE {json.dumps(msg)}")
|
|
||||||
await self._websocket.send(json.dumps(msg))
|
|
||||||
yield None
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"{self} exception: {e}")
|
logger.exception(f"{self} exception: {e}")
|
||||||
|
|||||||
Reference in New Issue
Block a user