Merge pull request #1603 from pipecat-ai/aleix/deepgram-tavus-fixes
deepgram/tavus fixes
This commit is contained in:
10
CHANGELOG.md
10
CHANGELOG.md
@@ -9,6 +9,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|
||||||
|
- `DeepgramTTSService` accepts `base_url` argument again, allowing you to
|
||||||
|
connect to an on-prem service.
|
||||||
|
|
||||||
- It is now possible to disable `SoundfileMixer` when created. You can then use
|
- It is now possible to disable `SoundfileMixer` when created. You can then use
|
||||||
`MixerEnableFrame` to dynamically enable it when necessary.
|
`MixerEnableFrame` to dynamically enable it when necessary.
|
||||||
|
|
||||||
@@ -30,8 +33,15 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
- `SoundfileMixer` constructor arguments need to be keywords.
|
- `SoundfileMixer` constructor arguments need to be keywords.
|
||||||
|
|
||||||
|
### Deprecated
|
||||||
|
|
||||||
|
- `DeepgramSTTService` parameter `url` is now deprecated, use `base_url`
|
||||||
|
instead.
|
||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
|
- Fixed a `TavusVideoService` issue that was causing audio choppiness.
|
||||||
|
|
||||||
- Fixed an issue in `SmallWebRTCTransport` where an error was thrown if the
|
- Fixed an issue in `SmallWebRTCTransport` where an error was thrown if the
|
||||||
client did not create a video transceiver.
|
client did not create a video transceiver.
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,6 @@
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
|
|
||||||
import aiohttp
|
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -40,104 +39,101 @@ async def run_bot(webrtc_connection: SmallWebRTCConnection):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Create an HTTP session
|
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
||||||
async with aiohttp.ClientSession() as session:
|
|
||||||
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
|
||||||
|
|
||||||
tts = DeepgramTTSService(
|
tts = DeepgramTTSService(
|
||||||
aiohttp_session=session,
|
api_key=os.getenv("DEEPGRAM_API_KEY"),
|
||||||
api_key=os.getenv("DEEPGRAM_API_KEY"),
|
voice="aura-asteria-en",
|
||||||
voice="aura-asteria-en",
|
base_url="http://0.0.0.0:8080",
|
||||||
base_url="http://0.0.0.0:8080/v1/speak",
|
)
|
||||||
)
|
|
||||||
|
|
||||||
llm = OpenAILLMService(
|
llm = OpenAILLMService(
|
||||||
# To use OpenAI
|
# To use OpenAI
|
||||||
# api_key=os.getenv("OPENAI_API_KEY"),
|
# api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
# Or, to use a local vLLM (or similar) api server
|
# Or, to use a local vLLM (or similar) api server
|
||||||
model="meta-llama/Meta-Llama-3-8B-Instruct",
|
model="meta-llama/Meta-Llama-3-8B-Instruct",
|
||||||
base_url="http://0.0.0.0:8000/v1",
|
base_url="http://0.0.0.0:8000/v1",
|
||||||
)
|
)
|
||||||
|
|
||||||
messages = [
|
messages = [
|
||||||
{
|
{
|
||||||
"role": "system",
|
"role": "system",
|
||||||
"content": "You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be converted to audio so don't include special characters in your answers. Respond to what the user said in a creative and helpful way.",
|
"content": "You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be converted to audio so don't include special characters in your answers. Respond to what the user said in a creative and helpful way.",
|
||||||
},
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
context = OpenAILLMContext(messages)
|
||||||
|
context_aggregator = llm.create_context_aggregator(context)
|
||||||
|
|
||||||
|
pipeline = Pipeline(
|
||||||
|
[
|
||||||
|
transport.input(), # Transport user input
|
||||||
|
stt, # STT
|
||||||
|
context_aggregator.user(),
|
||||||
|
llm, # LLM
|
||||||
|
tts, # TTS
|
||||||
|
transport.output(), # Transport bot output
|
||||||
|
context_aggregator.assistant(),
|
||||||
]
|
]
|
||||||
|
)
|
||||||
|
|
||||||
context = OpenAILLMContext(messages)
|
task = PipelineTask(
|
||||||
context_aggregator = llm.create_context_aggregator(context)
|
pipeline,
|
||||||
|
params=PipelineParams(
|
||||||
|
allow_interruptions=True,
|
||||||
|
enable_metrics=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
pipeline = Pipeline(
|
# When the first participant joins, the bot should introduce itself.
|
||||||
[
|
@transport.event_handler("on_client_connected")
|
||||||
transport.input(), # Transport user input
|
async def on_client_connected(transport, client):
|
||||||
stt, # STT
|
logger.info(f"Client connected")
|
||||||
context_aggregator.user(),
|
# Kick off the conversation.
|
||||||
llm, # LLM
|
messages.append({"role": "system", "content": "Please introduce yourself to the user."})
|
||||||
tts, # TTS
|
await task.queue_frames([context_aggregator.user().get_context_frame()])
|
||||||
transport.output(), # Transport bot output
|
|
||||||
context_aggregator.assistant(),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
|
|
||||||
task = PipelineTask(
|
# Handle "latency-ping" messages. The client will send app messages that look like
|
||||||
pipeline,
|
# this:
|
||||||
params=PipelineParams(
|
# { "latency-ping": { ts: <client-side timestamp> }}
|
||||||
allow_interruptions=True,
|
#
|
||||||
enable_metrics=True,
|
# We want to send an immediate pong back to the client from this handler function.
|
||||||
),
|
# Also, we will push a frame into the top of the pipeline and send it after the
|
||||||
)
|
#
|
||||||
|
@transport.event_handler("on_app_message")
|
||||||
# When the first participant joins, the bot should introduce itself.
|
async def on_app_message(transport, message, sender):
|
||||||
@transport.event_handler("on_client_connected")
|
try:
|
||||||
async def on_client_connected(transport, client):
|
if "latency-ping" in message:
|
||||||
logger.info(f"Client connected")
|
logger.debug(f"Received latency ping app message: {message}")
|
||||||
# Kick off the conversation.
|
ts = message["latency-ping"]["ts"]
|
||||||
messages.append({"role": "system", "content": "Please introduce yourself to the user."})
|
# Send immediately
|
||||||
await task.queue_frames([context_aggregator.user().get_context_frame()])
|
transport.output().send_message(
|
||||||
|
DailyTransportMessageFrame(
|
||||||
# Handle "latency-ping" messages. The client will send app messages that look like
|
message={"latency-pong-msg-handler": {"ts": ts}}, participant_id=sender
|
||||||
# this:
|
|
||||||
# { "latency-ping": { ts: <client-side timestamp> }}
|
|
||||||
#
|
|
||||||
# We want to send an immediate pong back to the client from this handler function.
|
|
||||||
# Also, we will push a frame into the top of the pipeline and send it after the
|
|
||||||
#
|
|
||||||
@transport.event_handler("on_app_message")
|
|
||||||
async def on_app_message(transport, message, sender):
|
|
||||||
try:
|
|
||||||
if "latency-ping" in message:
|
|
||||||
logger.debug(f"Received latency ping app message: {message}")
|
|
||||||
ts = message["latency-ping"]["ts"]
|
|
||||||
# Send immediately
|
|
||||||
transport.output().send_message(
|
|
||||||
DailyTransportMessageFrame(
|
|
||||||
message={"latency-pong-msg-handler": {"ts": ts}}, participant_id=sender
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
# And push to the pipeline for the Daily transport.output to send
|
)
|
||||||
await task.queue_frame(
|
# And push to the pipeline for the Daily transport.output to send
|
||||||
DailyTransportMessageFrame(
|
await task.queue_frame(
|
||||||
message={"latency-pong-pipeline-delivery": {"ts": ts}},
|
DailyTransportMessageFrame(
|
||||||
participant_id=sender,
|
message={"latency-pong-pipeline-delivery": {"ts": ts}},
|
||||||
)
|
participant_id=sender,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
)
|
||||||
logger.debug(f"message handling error: {e} - {message}")
|
except Exception as e:
|
||||||
|
logger.debug(f"message handling error: {e} - {message}")
|
||||||
|
|
||||||
@transport.event_handler("on_client_disconnected")
|
@transport.event_handler("on_client_disconnected")
|
||||||
async def on_client_disconnected(transport, client):
|
async def on_client_disconnected(transport, client):
|
||||||
logger.info(f"Client disconnected")
|
logger.info(f"Client disconnected")
|
||||||
|
|
||||||
@transport.event_handler("on_client_closed")
|
@transport.event_handler("on_client_closed")
|
||||||
async def on_client_closed(transport, client):
|
async def on_client_closed(transport, client):
|
||||||
logger.info(f"Client closed connection")
|
logger.info(f"Client closed connection")
|
||||||
await task.cancel()
|
await task.cancel()
|
||||||
|
|
||||||
runner = PipelineRunner(handle_sigint=False)
|
runner = PipelineRunner(handle_sigint=False)
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -231,9 +231,9 @@ class PollyTTSService(TTSService):
|
|||||||
|
|
||||||
yield TTSStartedFrame()
|
yield TTSStartedFrame()
|
||||||
|
|
||||||
chunk_size = 8192
|
CHUNK_SIZE = 1024
|
||||||
for i in range(0, len(audio_data), chunk_size):
|
for i in range(0, len(audio_data), CHUNK_SIZE):
|
||||||
chunk = audio_data[i : i + chunk_size]
|
chunk = audio_data[i : i + CHUNK_SIZE]
|
||||||
if len(chunk) > 0:
|
if len(chunk) > 0:
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
frame = TTSAudioRawFrame(chunk, self.sample_rate, 1)
|
frame = TTSAudioRawFrame(chunk, self.sample_rate, 1)
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ class DeepgramSTTService(STTService):
|
|||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
url: str = "",
|
url: str = "",
|
||||||
|
base_url: str = "",
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
live_options: Optional[LiveOptions] = None,
|
live_options: Optional[LiveOptions] = None,
|
||||||
addons: Optional[Dict] = None,
|
addons: Optional[Dict] = None,
|
||||||
@@ -53,6 +54,17 @@ class DeepgramSTTService(STTService):
|
|||||||
sample_rate = sample_rate or (live_options.sample_rate if live_options else None)
|
sample_rate = sample_rate or (live_options.sample_rate if live_options else None)
|
||||||
super().__init__(sample_rate=sample_rate, **kwargs)
|
super().__init__(sample_rate=sample_rate, **kwargs)
|
||||||
|
|
||||||
|
if url:
|
||||||
|
import warnings
|
||||||
|
|
||||||
|
with warnings.catch_warnings():
|
||||||
|
warnings.simplefilter("always")
|
||||||
|
warnings.warn(
|
||||||
|
"Parameter 'url' is deprecated, use 'base_url' instead.",
|
||||||
|
DeprecationWarning,
|
||||||
|
)
|
||||||
|
base_url = url
|
||||||
|
|
||||||
default_options = LiveOptions(
|
default_options = LiveOptions(
|
||||||
encoding="linear16",
|
encoding="linear16",
|
||||||
language=Language.EN,
|
language=Language.EN,
|
||||||
@@ -81,7 +93,7 @@ class DeepgramSTTService(STTService):
|
|||||||
self._client = DeepgramClient(
|
self._client = DeepgramClient(
|
||||||
api_key,
|
api_key,
|
||||||
config=DeepgramClientOptions(
|
config=DeepgramClientOptions(
|
||||||
url=url,
|
url=base_url,
|
||||||
options={"keepalive": "true"}, # verbose=logging.DEBUG
|
options={"keepalive": "true"}, # verbose=logging.DEBUG
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -4,7 +4,6 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
import asyncio
|
|
||||||
from typing import AsyncGenerator, Optional
|
from typing import AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -19,7 +18,7 @@ from pipecat.frames.frames import (
|
|||||||
from pipecat.services.tts_service import TTSService
|
from pipecat.services.tts_service import TTSService
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from deepgram import DeepgramClient, SpeakOptions
|
from deepgram import DeepgramClient, DeepgramClientOptions, SpeakOptions
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
logger.error(f"Exception: {e}")
|
logger.error(f"Exception: {e}")
|
||||||
logger.error("In order to use Deepgram, you need to `pip install pipecat-ai[deepgram]`.")
|
logger.error("In order to use Deepgram, you need to `pip install pipecat-ai[deepgram]`.")
|
||||||
@@ -32,6 +31,7 @@ class DeepgramTTSService(TTSService):
|
|||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
voice: str = "aura-helios-en",
|
voice: str = "aura-helios-en",
|
||||||
|
base_url: str = "",
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
encoding: str = "linear16",
|
encoding: str = "linear16",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
@@ -42,7 +42,9 @@ class DeepgramTTSService(TTSService):
|
|||||||
"encoding": encoding,
|
"encoding": encoding,
|
||||||
}
|
}
|
||||||
self.set_voice(voice)
|
self.set_voice(voice)
|
||||||
self._deepgram_client = DeepgramClient(api_key=api_key)
|
|
||||||
|
client_options = DeepgramClientOptions(url=base_url)
|
||||||
|
self._deepgram_client = DeepgramClient(api_key, config=client_options)
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -60,8 +62,8 @@ class DeepgramTTSService(TTSService):
|
|||||||
try:
|
try:
|
||||||
await self.start_ttfb_metrics()
|
await self.start_ttfb_metrics()
|
||||||
|
|
||||||
response = await asyncio.to_thread(
|
response = await self._deepgram_client.speak.asyncrest.v("1").stream_memory(
|
||||||
self._deepgram_client.speak.v("1").stream, {"text": text}, options
|
{"text": text}, options
|
||||||
)
|
)
|
||||||
|
|
||||||
await self.start_tts_usage_metrics(text)
|
await self.start_tts_usage_metrics(text)
|
||||||
|
|||||||
@@ -550,7 +550,7 @@ class ElevenLabsHttpTTSService(TTSService):
|
|||||||
if self._settings["optimize_streaming_latency"] is not None:
|
if self._settings["optimize_streaming_latency"] is not None:
|
||||||
params["optimize_streaming_latency"] = self._settings["optimize_streaming_latency"]
|
params["optimize_streaming_latency"] = self._settings["optimize_streaming_latency"]
|
||||||
|
|
||||||
logger.debug(f"ElevenLabs request - payload: {payload}, params: {params}")
|
logger.debug(f"{self} ElevenLabs request - payload: {payload}, params: {params}")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await self.start_ttfb_metrics()
|
await self.start_ttfb_metrics()
|
||||||
|
|||||||
@@ -346,9 +346,9 @@ class GoogleTTSService(TTSService):
|
|||||||
audio_content = response.audio_content[44:]
|
audio_content = response.audio_content[44:]
|
||||||
|
|
||||||
# Read and yield audio data in chunks
|
# Read and yield audio data in chunks
|
||||||
chunk_size = 8192
|
CHUNK_SIZE = 1024
|
||||||
for i in range(0, len(audio_content), chunk_size):
|
for i in range(0, len(audio_content), CHUNK_SIZE):
|
||||||
chunk = audio_content[i : i + chunk_size]
|
chunk = audio_content[i : i + CHUNK_SIZE]
|
||||||
if not chunk:
|
if not chunk:
|
||||||
break
|
break
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
|
|||||||
@@ -6,7 +6,9 @@
|
|||||||
|
|
||||||
"""This module implements Tavus as a sink transport layer"""
|
"""This module implements Tavus as a sink transport layer"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import base64
|
import base64
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -16,6 +18,7 @@ from pipecat.frames.frames import (
|
|||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
|
StartFrame,
|
||||||
StartInterruptionFrame,
|
StartInterruptionFrame,
|
||||||
TransportMessageUrgentFrame,
|
TransportMessageUrgentFrame,
|
||||||
TTSAudioRawFrame,
|
TTSAudioRawFrame,
|
||||||
@@ -50,6 +53,10 @@ class TavusVideoService(AIService):
|
|||||||
|
|
||||||
self._resampler = create_default_resampler()
|
self._resampler = create_default_resampler()
|
||||||
|
|
||||||
|
self._audio_buffer = bytearray()
|
||||||
|
self._queue = asyncio.Queue()
|
||||||
|
self._send_task: Optional[asyncio.Task] = None
|
||||||
|
|
||||||
async def initialize(self) -> str:
|
async def initialize(self) -> str:
|
||||||
url = "https://tavusapi.com/v2/conversations"
|
url = "https://tavusapi.com/v2/conversations"
|
||||||
headers = {"Content-Type": "application/json", "x-api-key": self._api_key}
|
headers = {"Content-Type": "application/json", "x-api-key": self._api_key}
|
||||||
@@ -78,45 +85,98 @@ class TavusVideoService(AIService):
|
|||||||
logger.debug(f"TavusVideoService persona grabbed {response_json}")
|
logger.debug(f"TavusVideoService persona grabbed {response_json}")
|
||||||
return response_json["persona_name"]
|
return response_json["persona_name"]
|
||||||
|
|
||||||
|
async def start(self, frame: StartFrame):
|
||||||
|
await super().start(frame)
|
||||||
|
await self._create_send_task()
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
await super().stop(frame)
|
await super().stop(frame)
|
||||||
await self._end_conversation()
|
await self._end_conversation()
|
||||||
|
await self._cancel_send_task()
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
await super().cancel(frame)
|
await super().cancel(frame)
|
||||||
await self._end_conversation()
|
await self._end_conversation()
|
||||||
|
await self._cancel_send_task()
|
||||||
|
|
||||||
async def _end_conversation(self) -> None:
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
|
if isinstance(frame, StartInterruptionFrame):
|
||||||
|
await self._handle_interruptions()
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
elif isinstance(frame, TTSStartedFrame):
|
||||||
|
await self.start_processing_metrics()
|
||||||
|
await self.start_ttfb_metrics()
|
||||||
|
self._current_idx_str = str(frame.id)
|
||||||
|
elif isinstance(frame, TTSAudioRawFrame):
|
||||||
|
await self._queue_audio(frame.audio, frame.sample_rate, done=False)
|
||||||
|
elif isinstance(frame, TTSStoppedFrame):
|
||||||
|
await self._queue_audio(b"\x00\x00", self._sample_rate, done=True)
|
||||||
|
await self.stop_ttfb_metrics()
|
||||||
|
await self.stop_processing_metrics()
|
||||||
|
else:
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
async def _handle_interruptions(self):
|
||||||
|
await self._cancel_send_task()
|
||||||
|
await self._create_send_task()
|
||||||
|
await self._send_interrupt_message()
|
||||||
|
|
||||||
|
async def _end_conversation(self):
|
||||||
url = f"https://tavusapi.com/v2/conversations/{self._conversation_id}/end"
|
url = f"https://tavusapi.com/v2/conversations/{self._conversation_id}/end"
|
||||||
headers = {"Content-Type": "application/json", "x-api-key": self._api_key}
|
headers = {"Content-Type": "application/json", "x-api-key": self._api_key}
|
||||||
async with self._session.post(url, headers=headers) as r:
|
async with self._session.post(url, headers=headers) as r:
|
||||||
r.raise_for_status()
|
r.raise_for_status()
|
||||||
|
|
||||||
async def _encode_audio_and_send(self, audio: bytes, in_rate: int, done: bool) -> None:
|
async def _queue_audio(self, audio: bytes, in_rate: int, done: bool):
|
||||||
|
await self._queue.put((audio, in_rate, done))
|
||||||
|
|
||||||
|
async def _create_send_task(self):
|
||||||
|
if not self._send_task:
|
||||||
|
self._queue = asyncio.Queue()
|
||||||
|
self._send_task = self.create_task(self._send_task_handler())
|
||||||
|
|
||||||
|
async def _cancel_send_task(self):
|
||||||
|
if self._send_task:
|
||||||
|
await self.cancel_task(self._send_task)
|
||||||
|
self._send_task = None
|
||||||
|
|
||||||
|
async def _send_task_handler(self):
|
||||||
|
# Daily app-messages have a 4kb limit and also a rate limit of 20
|
||||||
|
# messages per second. Below, we only consider the rate limit because 1
|
||||||
|
# second of a 24000 sample rate would be 48000 bytes (16-bit samples and
|
||||||
|
# 1 channel). So, that is 48000 / 20 = 2400, which is below the 4kb
|
||||||
|
# limit (even including base64 encoding). For a sample rate of 16000,
|
||||||
|
# that would be 32000 / 20 = 1600.
|
||||||
|
MAX_CHUNK_SIZE = int((self._sample_rate * 2) / 20)
|
||||||
|
SLEEP_TIME = 1 / 20
|
||||||
|
|
||||||
|
audio_buffer = bytearray()
|
||||||
|
while True:
|
||||||
|
(audio, in_rate, done) = await self._queue.get()
|
||||||
|
|
||||||
|
if done:
|
||||||
|
# Send any remaining audio.
|
||||||
|
if len(audio_buffer) > 0:
|
||||||
|
await self._encode_audio_and_send(bytes(audio_buffer), done)
|
||||||
|
await self._encode_audio_and_send(audio, done)
|
||||||
|
audio_buffer.clear()
|
||||||
|
else:
|
||||||
|
audio = await self._resampler.resample(audio, in_rate, self._sample_rate)
|
||||||
|
audio_buffer.extend(audio)
|
||||||
|
while len(audio_buffer) >= MAX_CHUNK_SIZE:
|
||||||
|
chunk = audio_buffer[:MAX_CHUNK_SIZE]
|
||||||
|
audio_buffer = audio_buffer[MAX_CHUNK_SIZE:]
|
||||||
|
await self._encode_audio_and_send(bytes(chunk), done)
|
||||||
|
await asyncio.sleep(SLEEP_TIME)
|
||||||
|
|
||||||
|
async def _encode_audio_and_send(self, audio: bytes, done: bool):
|
||||||
"""Encodes audio to base64 and sends it to Tavus"""
|
"""Encodes audio to base64 and sends it to Tavus"""
|
||||||
if not done:
|
|
||||||
audio = await self._resampler.resample(audio, in_rate, self._sample_rate)
|
|
||||||
audio_base64 = base64.b64encode(audio).decode("utf-8")
|
audio_base64 = base64.b64encode(audio).decode("utf-8")
|
||||||
logger.trace(f"{self}: sending {len(audio)} bytes")
|
logger.trace(f"{self}: sending {len(audio)} bytes")
|
||||||
await self._send_audio_message(audio_base64, done=done)
|
await self._send_audio_message(audio_base64, done=done)
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
|
||||||
await super().process_frame(frame, direction)
|
|
||||||
if isinstance(frame, TTSStartedFrame):
|
|
||||||
await self.start_processing_metrics()
|
|
||||||
await self.start_ttfb_metrics()
|
|
||||||
self._current_idx_str = str(frame.id)
|
|
||||||
elif isinstance(frame, TTSAudioRawFrame):
|
|
||||||
await self._encode_audio_and_send(frame.audio, frame.sample_rate, done=False)
|
|
||||||
elif isinstance(frame, TTSStoppedFrame):
|
|
||||||
await self._encode_audio_and_send(b"\x00", self._sample_rate, done=True)
|
|
||||||
await self.stop_ttfb_metrics()
|
|
||||||
await self.stop_processing_metrics()
|
|
||||||
elif isinstance(frame, StartInterruptionFrame):
|
|
||||||
await self._send_interrupt_message()
|
|
||||||
else:
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
|
|
||||||
async def _send_interrupt_message(self) -> None:
|
async def _send_interrupt_message(self) -> None:
|
||||||
transport_frame = TransportMessageUrgentFrame(
|
transport_frame = TransportMessageUrgentFrame(
|
||||||
message={
|
message={
|
||||||
@@ -127,7 +187,7 @@ class TavusVideoService(AIService):
|
|||||||
)
|
)
|
||||||
await self.push_frame(transport_frame)
|
await self.push_frame(transport_frame)
|
||||||
|
|
||||||
async def _send_audio_message(self, audio_base64: str, done: bool) -> None:
|
async def _send_audio_message(self, audio_base64: str, done: bool):
|
||||||
transport_frame = TransportMessageUrgentFrame(
|
transport_frame = TransportMessageUrgentFrame(
|
||||||
message={
|
message={
|
||||||
"message_type": "conversation",
|
"message_type": "conversation",
|
||||||
|
|||||||
Reference in New Issue
Block a user