694 lines
27 KiB
Python
694 lines
27 KiB
Python
import asyncio
|
||
import json
|
||
import websockets
|
||
import numpy as np
|
||
import sounddevice as sd
|
||
from typing import Optional, Callable, Dict, Any, List
|
||
from dataclasses import dataclass, field
|
||
import uuid
|
||
import copy
|
||
from protocols import (
|
||
EventType,
|
||
MsgType,
|
||
finish_connection,
|
||
finish_session,
|
||
receive_message,
|
||
start_connection,
|
||
start_session,
|
||
task_request,
|
||
wait_for_event,
|
||
)
|
||
|
||
# ------------------------------
|
||
# 配置常量(可根据需求调整)
|
||
# ------------------------------
|
||
DEFAULT_APPID = "7069844318"
|
||
DEFAULT_ACCESS_TOKEN = "osFMEJr20SSTWRql43cJlZkAOg7iwvxu"
|
||
DEFAULT_ENDPOINT = "wss://openspeech.bytedance.com/api/v3/tts/bidirection"
|
||
DEFAULT_VOICE_TYPE = "zh_female_gaolengyujie_emo_v2_mars_bigtts"
|
||
DEFAULT_ENCODING = "pcm"
|
||
DEFAULT_SAMPLE_RATE = 16000
|
||
|
||
|
||
@dataclass
|
||
class TTSRequest:
|
||
"""TTS请求对象(带唯一标识)"""
|
||
tts_text: str
|
||
voice_type: str = DEFAULT_VOICE_TYPE
|
||
encoding: str = DEFAULT_ENCODING
|
||
speed: float = 1.0 # 语速(字节跳动TTS支持,需服务端兼容)
|
||
stream: bool = True # 是否流式合成
|
||
request_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 唯一请求ID
|
||
session_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 会话ID(每个请求一个会话)
|
||
|
||
|
||
class ByteDanceTTSSocketClient:
|
||
"""字节跳动 TTS WebSocket 客户端(异步/流式/带任务队列)"""
|
||
|
||
def __init__(
|
||
self,
|
||
appid: str = DEFAULT_APPID,
|
||
access_token: str = DEFAULT_ACCESS_TOKEN,
|
||
endpoint: str = DEFAULT_ENDPOINT,
|
||
max_queue_size: int = 100
|
||
):
|
||
"""
|
||
初始化客户端
|
||
:param appid: 字节跳动APP ID
|
||
:param access_token: 访问令牌
|
||
:param endpoint: WebSocket 服务端地址
|
||
:param max_queue_size: 最大队列长度(防止内存溢出)
|
||
"""
|
||
# 基础配置
|
||
self.appid = appid
|
||
self.access_token = access_token
|
||
self.endpoint = endpoint
|
||
self.max_queue_size = max_queue_size
|
||
|
||
# WebSocket 连接状态
|
||
self.websocket: Optional[websockets.WebSocketClientProtocol] = None
|
||
self.is_connected = False
|
||
self.is_processing = False # 是否正在处理请求
|
||
self.logid: Optional[str] = None # 服务端返回的日志ID
|
||
|
||
# 异步任务队列(FIFO)
|
||
self.request_queue: asyncio.Queue[TTSRequest] = asyncio.Queue(maxsize=max_queue_size)
|
||
|
||
# 回调函数定义(所有回调都带request_id,方便关联请求)
|
||
self.on_task_enqueue: Callable[[str], None] = lambda req_id: None # 任务入队回调
|
||
self.on_start: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成开始回调
|
||
self.on_audio_chunk: Callable[[str, bytes], None] = lambda req_id, chunk: None # 音频块回调(原始字节)
|
||
self.on_end: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成结束回调
|
||
self.on_error: Callable[[str, str], None] = lambda req_id, msg: None # 错误回调
|
||
self.on_queue_full: Callable[[str], None] = lambda req_id: None # 队列满回调
|
||
|
||
# 音频播放相关(支持MP3格式直接播放)
|
||
self.play_stream: Optional[sd.OutputStream] = None
|
||
self.current_req_id: Optional[str] = None
|
||
self.enable_playback: bool = True # 是否启用实时播放
|
||
|
||
# 外部回调函数(返回完整结果)
|
||
self.external_callback: Optional[Callable[[str, Dict[str, Any]], None]] = None
|
||
|
||
# 存储每个请求的完整音频数据(原始字节)
|
||
self.audio_buffers: Dict[str, List[bytes]] = {}
|
||
|
||
def _get_resource_id(self, voice_type: str) -> str:
|
||
"""根据音色类型获取资源ID(字节跳动TTS协议要求)"""
|
||
if voice_type.startswith("S_"):
|
||
return "volc.megatts.default"
|
||
return "volc.service_type.10029"
|
||
|
||
async def _create_websocket_connection(self):
|
||
"""创建WebSocket连接(内部使用)"""
|
||
headers = {
|
||
"X-Api-App-Key": self.appid,
|
||
"X-Api-Access-Key": self.access_token,
|
||
"X-Api-Resource-Id": self._get_resource_id(DEFAULT_VOICE_TYPE), # 用默认音色获取资源ID
|
||
"X-Api-Connect-Id": str(uuid.uuid4()),
|
||
}
|
||
|
||
print(f"连接到 TTS 服务端: {self.endpoint}")
|
||
self.websocket = await websockets.connect(
|
||
self.endpoint,
|
||
additional_headers=headers,
|
||
max_size=10 * 1024 * 1024 # 10MB缓冲区
|
||
)
|
||
self.is_connected = True
|
||
self.logid = self.websocket.response.headers.get("x-tt-logid")
|
||
print(f"连接成功,LogID: {self.logid}")
|
||
|
||
# 发送连接启动指令
|
||
await start_connection(self.websocket)
|
||
await wait_for_event(
|
||
self.websocket, MsgType.FullServerResponse, EventType.ConnectionStarted
|
||
)
|
||
print("TTS连接已初始化完成")
|
||
|
||
async def connect(self):
|
||
"""建立WebSocket连接(外部调用,初始化一次)"""
|
||
if not self.is_connected:
|
||
try:
|
||
await self._create_websocket_connection()
|
||
# 启动队列消费协程(后台运行)
|
||
asyncio.create_task(self._consume_queue())
|
||
except Exception as e:
|
||
error_msg = f"连接失败: {str(e)}"
|
||
print(error_msg)
|
||
raise ConnectionError(error_msg)
|
||
|
||
async def disconnect(self):
|
||
"""关闭WebSocket连接"""
|
||
if self.is_connected and self.websocket:
|
||
try:
|
||
# 发送连接结束指令
|
||
await finish_connection(self.websocket)
|
||
await wait_for_event(
|
||
self.websocket, MsgType.FullServerResponse, EventType.ConnectionFinished
|
||
)
|
||
except Exception as e:
|
||
print(f"关闭连接时异常: {str(e)}")
|
||
finally:
|
||
await self.websocket.close()
|
||
self.is_connected = False
|
||
self.websocket = None
|
||
print("已断开与TTS服务端的连接")
|
||
|
||
# 清理音频播放流
|
||
if self.play_stream:
|
||
self.play_stream.stop()
|
||
self.play_stream.close()
|
||
self.play_stream = None
|
||
|
||
def set_external_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
|
||
"""设置外部回调函数,用于返回完整结果"""
|
||
self.external_callback = callback
|
||
|
||
def set_playback_enabled(self, enabled: bool):
|
||
"""设置是否启用音频实时播放"""
|
||
self.enable_playback = enabled
|
||
print(f"音频实时播放已{'启用' if enabled else '禁用'}")
|
||
|
||
async def synthesize(self, tts_text: str, **kwargs) -> str:
|
||
"""
|
||
异步非阻塞添加TTS请求到队列
|
||
:param tts_text: 要合成的文本
|
||
:param kwargs: 其他TTS参数(voice_type, encoding, speed等)
|
||
:return: 唯一请求ID
|
||
"""
|
||
# 创建请求对象(支持覆盖默认参数)
|
||
request = TTSRequest(tts_text=tts_text, **kwargs)
|
||
req_id = request.request_id
|
||
|
||
# 初始化音频缓冲区
|
||
self.audio_buffers[req_id] = []
|
||
|
||
# 异步入队(非阻塞)
|
||
try:
|
||
await self.request_queue.put(request)
|
||
self.on_task_enqueue(req_id)
|
||
print(f"请求 [{req_id[:8]}] 已加入队列,当前队列长度: {self.request_queue.qsize()}")
|
||
return req_id
|
||
except asyncio.QueueFull:
|
||
self.on_queue_full(req_id)
|
||
error_msg = f"队列已满(最大长度{self.max_queue_size}),请求 [{req_id[:8]}] 入队失败"
|
||
print(error_msg)
|
||
raise Exception(error_msg)
|
||
|
||
async def _consume_queue(self):
|
||
"""消费队列(后台协程,自动处理排队请求)"""
|
||
print("队列消费协程已启动")
|
||
while True:
|
||
try:
|
||
# 等待队列中有请求(阻塞,直到有任务)
|
||
request = await self.request_queue.get()
|
||
req_id = request.request_id
|
||
|
||
# 标记为处理中
|
||
self.is_processing = True
|
||
# print(f"\n开始处理请求 [{req_id[:8]}],剩余队列长度: {self.request_queue.qsize()}")
|
||
|
||
# 处理单个请求
|
||
await self._process_single_request(request)
|
||
|
||
# 标记任务完成(让Queue知道可以继续)
|
||
self.request_queue.task_done()
|
||
self.is_processing = False
|
||
|
||
except Exception as e:
|
||
error_msg = f"队列消费异常: {str(e)}"
|
||
print(error_msg)
|
||
self.is_processing = False
|
||
# 短暂等待,避免死循环占用CPU
|
||
await asyncio.sleep(0.1)
|
||
|
||
def _build_base_request(self, request: TTSRequest) -> Dict[str, Any]:
|
||
"""构建字节跳动TTS基础请求参数"""
|
||
aaa = {
|
||
"user": {"uid": str(uuid.uuid4())},
|
||
"namespace": "BidirectionalTTS",
|
||
"req_params": {
|
||
"speaker": request.voice_type,
|
||
"audio_params": {
|
||
"format": request.encoding,
|
||
"sample_rate": DEFAULT_SAMPLE_RATE,
|
||
"enable_timestamp": True,
|
||
},
|
||
"additions": json.dumps({"disable_markdown_filter": False}),
|
||
"speed": request.speed, # 语速参数(需服务端支持)
|
||
},
|
||
}
|
||
print('aaa', aaa)
|
||
return aaa
|
||
async def _send_text_stream(self, request: TTSRequest, session_id: str):
|
||
"""流式发送文本(逐字符发送,字节跳动TTS流式协议要求)"""
|
||
base_request = self._build_base_request(request)
|
||
text = request.tts_text.strip()
|
||
|
||
if not text:
|
||
print(f"请求 [{request.request_id[:8]}] 文本为空,跳过发送")
|
||
return
|
||
|
||
# 逐字符发送(控制发送速率,避免拥塞)
|
||
for char in text:
|
||
if not self.is_connected or not self.websocket:
|
||
raise ConnectionError("连接已断开,无法继续发送文本")
|
||
|
||
# 构建单个字符的任务请求
|
||
task_req = copy.deepcopy(base_request)
|
||
task_req["event"] = EventType.TaskRequest
|
||
task_req["req_params"]["text"] = char
|
||
|
||
# 发送任务请求
|
||
await task_request(
|
||
self.websocket,
|
||
json.dumps(task_req).encode("utf-8"),
|
||
session_id
|
||
)
|
||
# 控制发送速率(5ms/字符,可调整)
|
||
await asyncio.sleep(0.005)
|
||
|
||
# 发送会话结束指令
|
||
await finish_session(self.websocket, session_id)
|
||
print(f"请求 [{request.request_id[:8]}] 文本发送完成")
|
||
|
||
async def _handle_audio_response(self, req_id: str, session_id: str, request: TTSRequest) -> Dict[str, Any]:
|
||
"""处理服务端的流式音频响应"""
|
||
if not self.websocket:
|
||
raise ConnectionError("WebSocket连接未建立")
|
||
|
||
sample_rate = DEFAULT_SAMPLE_RATE
|
||
audio_received = False
|
||
|
||
try:
|
||
while True:
|
||
# 接收服务端消息(异步阻塞)
|
||
msg = await receive_message(self.websocket)
|
||
|
||
if msg.type == MsgType.FullServerResponse:
|
||
# 完整响应(开始/结束/错误)
|
||
if msg.event == EventType.SessionStarted:
|
||
# 会话开始回调
|
||
start_data = {
|
||
"session_id": session_id,
|
||
"sample_rate": sample_rate,
|
||
"encoding": request.encoding,
|
||
"voice_type": request.voice_type,
|
||
"logid": self.logid
|
||
}
|
||
self.on_start(req_id, start_data)
|
||
print(f"请求 [{req_id[:8]}] 合成开始")
|
||
|
||
elif msg.event == EventType.SessionFinished:
|
||
# 会话结束,退出循环
|
||
end_data = {"session_id": session_id, "message": "合成完成"}
|
||
self.on_end(req_id, end_data)
|
||
print(f"请求 [{req_id[:8]}] 合成结束")
|
||
break
|
||
|
||
elif msg.type == MsgType.AudioOnlyServer:
|
||
# 流式音频数据(原始字节)
|
||
audio_chunk = msg.payload
|
||
if audio_chunk:
|
||
audio_received = True
|
||
# 保存到缓冲区
|
||
self.audio_buffers[req_id].append(audio_chunk)
|
||
# 音频块回调
|
||
self.on_audio_chunk(req_id, audio_chunk)
|
||
# 实时播放(如果启用)
|
||
await self._play_audio_chunk(req_id, audio_chunk, sample_rate)
|
||
|
||
else:
|
||
# 未知消息类型
|
||
raise RuntimeError(f"收到未知消息类型: {msg.type}, 内容: {msg}")
|
||
|
||
# 组装完整结果
|
||
full_audio = b"".join(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else b""
|
||
return {
|
||
"status": "completed",
|
||
"request_id": req_id,
|
||
"session_id": session_id,
|
||
"sample_rate": sample_rate,
|
||
"encoding": request.encoding,
|
||
"audio_data": full_audio, # 完整音频字节数据
|
||
"audio_length": len(full_audio),
|
||
"message": "合成成功" if audio_received else "合成完成但未收到音频数据"
|
||
}
|
||
|
||
except Exception as e:
|
||
error_msg = f"处理音频响应异常: {str(e)}"
|
||
self.on_error(req_id, error_msg)
|
||
return {
|
||
"status": "error",
|
||
"request_id": req_id,
|
||
"session_id": session_id,
|
||
"message": error_msg
|
||
}
|
||
|
||
async def _play_audio_chunk(self, req_id: str, chunk: bytes, sample_rate: int):
|
||
"""实时播放音频块(支持MP3格式)"""
|
||
if not self.enable_playback:
|
||
return
|
||
|
||
# 确保当前请求是正在播放的请求
|
||
if self.current_req_id is None:
|
||
self.current_req_id = req_id
|
||
|
||
if req_id != self.current_req_id:
|
||
# 切换请求时,重置播放流
|
||
if self.play_stream:
|
||
self.play_stream.stop()
|
||
self.play_stream.close()
|
||
self.current_req_id = req_id
|
||
|
||
try:
|
||
# 初始化播放流(如果未初始化)
|
||
if not self.play_stream:
|
||
self.play_stream = sd.OutputStream(
|
||
samplerate=sample_rate,
|
||
channels=1, # 单声道
|
||
dtype=np.float32
|
||
)
|
||
self.play_stream.start()
|
||
|
||
# MP3字节 → 音频数组(直接播放)
|
||
# 注意:sounddevice默认支持PCM格式,如果是MP3需要解码,这里简化处理(实际使用建议用pydub解码)
|
||
# 如需支持MP3播放,请安装 pydub: pip install pydub ffmpeg
|
||
try:
|
||
# 简化处理:假设服务端返回PCM(如果是MP3,需替换为解码逻辑)
|
||
audio_array = np.frombuffer(chunk, dtype=np.float32)
|
||
if audio_array.size > 0:
|
||
self.play_stream.write(audio_array)
|
||
except Exception as e:
|
||
print(f"音频播放异常: {str(e)},请确保音频格式正确")
|
||
|
||
except Exception as e:
|
||
print(f"播放流初始化失败: {str(e)}")
|
||
|
||
async def _process_single_request(self, request: TTSRequest):
|
||
"""处理单个TTS请求(完整流程:连接→启动会话→流式发送文本→接收音频→回调结果)"""
|
||
req_id = request.request_id
|
||
session_id = request.session_id
|
||
|
||
# 参数校验
|
||
if not request.tts_text.strip():
|
||
error_msg = "合成文本不能为空"
|
||
self.on_error(req_id, error_msg)
|
||
self._send_external_callback(req_id, {
|
||
"status": "error",
|
||
"request_id": req_id,
|
||
"session_id": session_id,
|
||
"message": error_msg
|
||
})
|
||
return
|
||
|
||
# 确保连接已建立(断开时自动重连)
|
||
if not self.is_connected:
|
||
print(f"请求 [{req_id[:8]}] 处理时连接已断开,尝试重连...")
|
||
try:
|
||
await self._create_websocket_connection()
|
||
except Exception as e:
|
||
error_msg = f"重连失败: {str(e)}"
|
||
self.on_error(req_id, error_msg)
|
||
self._send_external_callback(req_id, {
|
||
"status": "error",
|
||
"request_id": req_id,
|
||
"session_id": session_id,
|
||
"message": error_msg
|
||
})
|
||
return
|
||
|
||
try:
|
||
# 1. 启动会话
|
||
base_request = self._build_base_request(request)
|
||
start_session_req = copy.deepcopy(base_request)
|
||
start_session_req["event"] = EventType.StartSession
|
||
|
||
await start_session(
|
||
self.websocket,
|
||
json.dumps(start_session_req).encode("utf-8"),
|
||
session_id
|
||
)
|
||
# 等待会话启动成功
|
||
await wait_for_event(
|
||
self.websocket, MsgType.FullServerResponse, EventType.SessionStarted
|
||
)
|
||
|
||
# 2. 异步流式发送文本(后台任务,不阻塞接收音频)
|
||
send_task = asyncio.create_task(self._send_text_stream(request, session_id))
|
||
|
||
# 3. 接收并处理音频响应
|
||
result_data = await self._handle_audio_response(req_id, session_id, request)
|
||
|
||
# 4. 等待文本发送任务完成
|
||
await send_task
|
||
|
||
# 5. 发送外部回调
|
||
self._send_external_callback(req_id, result_data)
|
||
|
||
except Exception as e:
|
||
error_msg = f"处理请求 [{req_id[:8]}] 异常: {str(e)}"
|
||
self.on_error(req_id, error_msg)
|
||
self._send_external_callback(req_id, {
|
||
"status": "error",
|
||
"request_id": req_id,
|
||
"session_id": session_id,
|
||
"message": error_msg
|
||
})
|
||
finally:
|
||
# 清理缓冲区
|
||
if req_id in self.audio_buffers:
|
||
del self.audio_buffers[req_id]
|
||
|
||
def _send_external_callback(self, req_id: str, result_data: Dict[str, Any]):
|
||
"""发送外部回调(支持同步/异步回调函数)"""
|
||
if not self.external_callback:
|
||
return
|
||
|
||
try:
|
||
# 异步回调:直接await
|
||
if asyncio.iscoroutinefunction(self.external_callback):
|
||
asyncio.create_task(self.external_callback(req_id, result_data))
|
||
# 同步回调:在线程池中执行(避免阻塞事件循环)
|
||
else:
|
||
asyncio.get_event_loop().run_in_executor(
|
||
None, self.external_callback, req_id, result_data
|
||
)
|
||
except Exception as e:
|
||
print(f"外部回调执行异常: {str(e)}")
|
||
|
||
async def wait_all_completed(self):
|
||
"""等待队列中所有任务处理完成(阻塞)"""
|
||
await self.request_queue.join()
|
||
print("\n所有队列任务已处理完成")
|
||
|
||
|
||
# ------------------------------
|
||
# 使用示例(与你提供的风格完全一致)
|
||
# ------------------------------
|
||
class TTSManager:
|
||
"""TTS管理器 - 供外部代码调用(封装客户端,简化使用)"""
|
||
|
||
def __init__(
|
||
self,
|
||
appid: str = DEFAULT_APPID,
|
||
access_token: str = DEFAULT_ACCESS_TOKEN,
|
||
endpoint: str = DEFAULT_ENDPOINT
|
||
):
|
||
self.client = ByteDanceTTSSocketClient(
|
||
appid=appid,
|
||
access_token=access_token,
|
||
endpoint=endpoint
|
||
)
|
||
self._setup_internal_callbacks()
|
||
|
||
def _setup_internal_callbacks(self):
|
||
"""设置内部回调(日志/状态提示)"""
|
||
|
||
def on_task_enqueue(req_id: str):
|
||
"""任务入队回调"""
|
||
print(f"📥 任务 [{req_id[:8]}] 已入队")
|
||
|
||
def on_tts_start(req_id: str, data: Dict[str, Any]):
|
||
"""合成开始回调"""
|
||
print(f"🎤 合成开始 [{req_id[:8]}] - 采样率: {data['sample_rate']}, 编码: {data['encoding']}")
|
||
|
||
def on_audio_chunk(req_id: str, chunk: bytes):
|
||
"""音频块回调(内部仅打印日志,外部通过external_callback获取)"""
|
||
print(f"🔊 收到音频块 [{req_id[:8]}] - 大小: {len(chunk)}字节", end="\r")
|
||
|
||
def on_tts_end(req_id: str, data: Dict[str, Any]):
|
||
"""合成结束回调"""
|
||
print(f"\n🏁 合成结束 [{req_id[:8]}] - 会话ID: {data['session_id']}")
|
||
|
||
def on_tts_error(req_id: str, msg: str):
|
||
"""错误回调"""
|
||
print(f"\n❌ 合成失败 [{req_id[:8]}] - 错误: {msg}")
|
||
|
||
def on_queue_full(req_id: str):
|
||
"""队列满回调"""
|
||
print(f"⚠️ 队列已满,请求 [{req_id[:8]}] 入队失败")
|
||
|
||
# 绑定内部回调
|
||
self.client.on_task_enqueue = on_task_enqueue
|
||
self.client.on_start = on_tts_start
|
||
self.client.on_audio_chunk = on_audio_chunk
|
||
self.client.on_end = on_tts_end
|
||
self.client.on_error = on_tts_error
|
||
self.client.on_queue_full = on_queue_full
|
||
|
||
async def initialize(self):
|
||
"""初始化连接"""
|
||
await self.client.connect()
|
||
|
||
async def shutdown(self):
|
||
"""关闭连接"""
|
||
await self.client.disconnect()
|
||
|
||
def set_result_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
|
||
"""设置外部结果回调(获取完整音频数据)"""
|
||
self.client.set_external_callback(callback)
|
||
|
||
def set_playback_enabled(self, enabled: bool):
|
||
"""设置是否启用实时播放"""
|
||
self.client.set_playback_enabled(enabled)
|
||
|
||
async def synthesize(self, text: str, **kwargs) -> str:
|
||
"""
|
||
异步非阻塞合成文本
|
||
:param text: 要合成的文本
|
||
:param kwargs: 其他参数(voice_type, encoding, speed等)
|
||
:return: 请求ID
|
||
"""
|
||
return await self.client.synthesize(text, **kwargs)
|
||
|
||
async def wait_all_completed(self):
|
||
"""等待所有任务完成"""
|
||
await self.client.wait_all_completed()
|
||
|
||
|
||
# ------------------------------
|
||
# 外部调用示例
|
||
# ------------------------------
|
||
async def external_usage_example():
|
||
"""外部代码使用示例"""
|
||
# 1. 创建TTS管理器(可替换为自己的appid和access_token)
|
||
tts_manager = TTSManager(
|
||
appid=DEFAULT_APPID,
|
||
access_token=DEFAULT_ACCESS_TOKEN,
|
||
endpoint=DEFAULT_ENDPOINT
|
||
)
|
||
|
||
# 2. 设置外部结果回调(获取完整音频数据)
|
||
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
||
"""处理TTS完整结果(同步回调)"""
|
||
status = result.get("status")
|
||
if status == "completed":
|
||
audio_data = result.get("audio_data")
|
||
encoding = result.get("encoding")
|
||
audio_length = result.get("audio_length")
|
||
|
||
print(f"\n✅ 收到完整结果 [{req_id[:8]}] - 长度: {audio_length}字节, 编码: {encoding}")
|
||
|
||
# 保存音频文件
|
||
filename = f"tts_output_{req_id[:8]}.{encoding}"
|
||
with open(filename, "wb") as f:
|
||
f.write(audio_data)
|
||
print(f"💾 音频文件已保存: {filename}")
|
||
|
||
elif status == "error":
|
||
error_msg = result.get("message")
|
||
print(f"\n❌ 请求 [{req_id[:8]}] 处理失败: {error_msg}")
|
||
|
||
# 绑定外部回调
|
||
tts_manager.set_result_callback(handle_tts_result)
|
||
|
||
# 3. 设置是否启用实时播放(默认True)
|
||
tts_manager.set_playback_enabled(True)
|
||
|
||
# 4. 初始化连接
|
||
await tts_manager.initialize()
|
||
|
||
# 5. 异步提交多个TTS请求(非阻塞)
|
||
texts = [
|
||
"你好,这是字节跳动TTS的流式合成测试。",
|
||
"我支持异步非阻塞调用,多个请求可以排队处理。",
|
||
"每个请求都会返回唯一的ID,方便你跟踪结果。",
|
||
"音频数据会通过回调函数返回,支持实时播放和保存文件。",
|
||
"最后一个测试句子,演示队列的自动消费功能。"
|
||
]
|
||
|
||
req_ids = []
|
||
for i, text in enumerate(texts):
|
||
# 提交请求(非阻塞,立即返回)
|
||
req_id = await tts_manager.synthesize(
|
||
text,
|
||
voice_type=DEFAULT_VOICE_TYPE,
|
||
encoding=DEFAULT_ENCODING,
|
||
speed=1.0
|
||
)
|
||
req_ids.append(req_id)
|
||
print(f"📤 已提交请求 {i+1}: ID={req_id[:8]}")
|
||
|
||
# 模拟其他业务逻辑(无需等待TTS完成)
|
||
await asyncio.sleep(0.3)
|
||
|
||
# 6. 等待所有TTS任务完成(可选,根据业务需求决定是否等待)
|
||
await tts_manager.wait_all_completed()
|
||
|
||
# 7. 关闭连接(程序退出前调用)
|
||
await tts_manager.shutdown()
|
||
|
||
|
||
# ------------------------------
|
||
# 异步结果回调示例(高级用法)
|
||
# ------------------------------
|
||
async def async_result_callback(req_id: str, result: Dict[str, Any]):
|
||
"""异步结果回调(支持异步操作,如上传音频到服务器)"""
|
||
if result["status"] == "completed":
|
||
print(f"\n⚡ 异步处理结果 [{req_id[:8]}] - 开始上传音频...")
|
||
# 模拟异步上传操作
|
||
await asyncio.sleep(0.5)
|
||
print(f"⚡ 异步处理结果 [{req_id[:8]}] - 音频上传完成")
|
||
|
||
|
||
async def advanced_usage_example():
|
||
"""高级使用示例:异步回调 + 禁用播放 + 批量请求"""
|
||
tts_manager = TTSManager()
|
||
|
||
# 设置异步结果回调
|
||
tts_manager.set_result_callback(async_result_callback)
|
||
|
||
# 禁用实时播放(只获取音频数据)
|
||
tts_manager.set_playback_enabled(False)
|
||
|
||
await tts_manager.initialize()
|
||
|
||
# 批量提交请求(并行提交)
|
||
tasks = []
|
||
for i in range(3):
|
||
text = f"这是第{i+1}个高级测试文本,使用异步回调处理结果。"
|
||
task = tts_manager.synthesize(text, speed=0.9)
|
||
tasks.append(task)
|
||
|
||
# 并行提交所有请求
|
||
req_ids = await asyncio.gather(*tasks)
|
||
print(f"\n已并行提交 {len(req_ids)} 个请求")
|
||
|
||
# 等待所有任务完成
|
||
await tts_manager.wait_all_completed()
|
||
await tts_manager.shutdown()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
try:
|
||
# 运行基础使用示例
|
||
asyncio.run(external_usage_example())
|
||
|
||
# 运行高级使用示例(取消注释)
|
||
# asyncio.run(advanced_usage_example())
|
||
|
||
except KeyboardInterrupt:
|
||
print("\n程序被用户中断")
|
||
except Exception as e:
|
||
print(f"程序异常: {str(e)}") |