Files
2025-12-04 02:16:51 +08:00

693 lines
27 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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, # 语速参数(需服务端支持)
},
}
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)}")