MiniMaxHttpTTSService: renamed from Hailuo and move to services/minimax/tts.py

This commit is contained in:
Aleix Conchillo Flaqué
2025-05-15 11:58:54 -07:00
parent a51af35024
commit 3933ba57b8
2 changed files with 50 additions and 37 deletions

View File

@@ -9,6 +9,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
### Added ### Added
- Added support for a new TTS service, `MiniMaxHttpTTSService`.
(see https://www.minimax.io/audio)
- A new function `FrameProcessor.setup()` has been added to allow setting up - A new function `FrameProcessor.setup()` has been added to allow setting up
frame processors before receiving a `StartFrame`. This is what's happening frame processors before receiving a `StartFrame`. This is what's happening
internally: `FrameProcessor.setup()` is called, `StartFrame` is pushed from internally: `FrameProcessor.setup()` is called, `StartFrame` is pushed from

View File

@@ -1,21 +1,31 @@
from typing import AsyncGenerator, Optional #
import aiohttp # Copyright (c) 20242025, Daily
#
# SPDX-License-Identifier: BSD 2-Clause License
#
import json import json
import time
from typing import AsyncGenerator, Optional
import aiohttp
from loguru import logger from loguru import logger
from pydantic import BaseModel from pydantic import BaseModel
import time
from pipecat.frames.frames import ( from pipecat.frames.frames import (
CancelFrame,
EndFrame,
ErrorFrame, ErrorFrame,
Frame, Frame,
StartFrame,
TTSAudioRawFrame, TTSAudioRawFrame,
TTSStartedFrame, TTSStartedFrame,
TTSStoppedFrame, TTSStoppedFrame,
) )
from pipecat.services.ai_services import TTSService from pipecat.services.ai_services import TTSService
from pipecat.transcriptions.language import Language
class HailuoHttpTTSService(TTSService):
class MiniMaxHttpTTSService(TTSService):
class InputParams(BaseModel): class InputParams(BaseModel):
speed: Optional[float] = 1.0 speed: Optional[float] = 1.0
volume: Optional[float] = 1.0 volume: Optional[float] = 1.0
@@ -46,14 +56,14 @@ class HailuoHttpTTSService(TTSService):
"voice_id": voice_id, "voice_id": voice_id,
"speed": params.speed, "speed": params.speed,
"vol": params.volume, "vol": params.volume,
"pitch": params.pitch "pitch": params.pitch,
}, },
"audio_setting": { "audio_setting": {
"sample_rate": sample_rate, "sample_rate": sample_rate,
"bitrate": 128000, "bitrate": 128000,
"format": "pcm", "format": "pcm",
"channel": 1 "channel": 1,
} },
} }
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
@@ -68,21 +78,21 @@ class HailuoHttpTTSService(TTSService):
await self._session.close() await self._session.close()
self._session = None self._session = None
async def start(self, frame: Frame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._init_session() await self._init_session()
async def stop(self, frame: Frame): async def stop(self, frame: EndFrame):
await super().stop(frame) await super().stop(frame)
await self._close_session() await self._close_session()
async def cancel(self, frame: Frame): async def cancel(self, frame: CancelFrame):
await super().cancel(frame) await super().cancel(frame)
await self._close_session() await self._close_session()
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
text = text.strip() text = text.strip()
if not text or text in ['"', "'", ']', '[']: if not text or text in ['"', "'", "]", "["]:
logger.debug(f"Skipping invalid text for TTS: [{text}]") logger.debug(f"Skipping invalid text for TTS: [{text}]")
return return
@@ -95,9 +105,9 @@ class HailuoHttpTTSService(TTSService):
yield TTSStartedFrame() yield TTSStartedFrame()
headers = { headers = {
'accept': 'application/json, text/plain, */*', "accept": "application/json, text/plain, */*",
'Content-Type': 'application/json', "Content-Type": "application/json",
'Authorization': f'Bearer {self._api_key}' "Authorization": f"Bearer {self._api_key}",
} }
payload = { payload = {
@@ -105,13 +115,11 @@ class HailuoHttpTTSService(TTSService):
"text": text, "text": text,
"stream": True, "stream": True,
"voice_setting": self._settings["voice_setting"], "voice_setting": self._settings["voice_setting"],
"audio_setting": self._settings["audio_setting"] "audio_setting": self._settings["audio_setting"],
} }
async with self._session.post( async with self._session.post(
self._base_url, self._base_url, headers=headers, json=payload
headers=headers,
json=payload
) as response: ) as response:
if response.status != 200: if response.status != 200:
error_text = await response.text() error_text = await response.text()
@@ -129,9 +137,9 @@ class HailuoHttpTTSService(TTSService):
buffer.extend(chunk) buffer.extend(chunk)
# Find complete data blocks # Find complete data blocks
while b'data:' in buffer: while b"data:" in buffer:
start = buffer.find(b'data:') start = buffer.find(b"data:")
next_start = buffer.find(b'data:', start + 5) next_start = buffer.find(b"data:", start + 5)
if next_start == -1: if next_start == -1:
# No next data block found, keep current data for next iteration # No next data block found, keep current data for next iteration
@@ -144,7 +152,7 @@ class HailuoHttpTTSService(TTSService):
buffer = buffer[next_start:] buffer = buffer[next_start:]
try: try:
data = json.loads(data_block[5:].decode('utf-8')) data = json.loads(data_block[5:].decode("utf-8"))
# Skip data blocks containing extra_info # Skip data blocks containing extra_info
if "extra_info" in data: if "extra_info" in data:
logger.debug("Received final chunk with extra info") logger.debug("Received final chunk with extra info")
@@ -162,7 +170,7 @@ class HailuoHttpTTSService(TTSService):
CHUNK_SIZE = 4096 # 4KB per chunk CHUNK_SIZE = 4096 # 4KB per chunk
for i in range(0, len(audio_data), CHUNK_SIZE * 2): # *2 for hex string for i in range(0, len(audio_data), CHUNK_SIZE * 2): # *2 for hex string
# Split hex string # Split hex string
hex_chunk = audio_data[i:i + CHUNK_SIZE * 2] hex_chunk = audio_data[i : i + CHUNK_SIZE * 2]
if not hex_chunk: if not hex_chunk:
continue continue
@@ -172,8 +180,10 @@ class HailuoHttpTTSService(TTSService):
if audio_chunk: if audio_chunk:
yield TTSAudioRawFrame( yield TTSAudioRawFrame(
audio=audio_chunk, audio=audio_chunk,
sample_rate=self._settings["audio_setting"]["sample_rate"], sample_rate=self._settings["audio_setting"][
num_channels=self._settings["audio_setting"]["channel"] "sample_rate"
],
num_channels=self._settings["audio_setting"]["channel"],
) )
except ValueError as e: except ValueError as e:
logger.error(f"Error converting hex to binary: {e}") logger.error(f"Error converting hex to binary: {e}")