551 lines
20 KiB
Python
551 lines
20 KiB
Python
import asyncio
|
||
import json
|
||
import websockets
|
||
import numpy as np
|
||
import sounddevice as sd
|
||
from typing import Optional, Callable, Dict, Any, Coroutine, List
|
||
from dataclasses import dataclass, field
|
||
import uuid
|
||
|
||
|
||
@dataclass
|
||
class TTSRequest:
|
||
"""TTS请求对象(带唯一标识)"""
|
||
tts_text: str
|
||
mode: str = "预训练音色"
|
||
sft_spk: str = ""
|
||
seed: int = field(default_factory=lambda: np.random.randint(1, 100000000))
|
||
stream: bool = True
|
||
speed: float = 1.0
|
||
prompt_wav: str = ""
|
||
instruct_text: str = ""
|
||
request_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 唯一请求ID
|
||
|
||
|
||
class CosyVoiceTTSSocketClient:
|
||
"""CosyVoice TTS WebSocket 客户端(异步/流式/带任务队列)"""
|
||
|
||
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts", max_queue_size: int = 100):
|
||
"""
|
||
初始化客户端
|
||
:param ws_url: WebSocket 服务端地址
|
||
:param max_queue_size: 最大队列长度(防止内存溢出)
|
||
"""
|
||
self.ws_url = ws_url
|
||
self.websocket: Optional[websockets.WebSocketClientProtocol] = None
|
||
self.is_connected = False
|
||
self.is_processing = False # 是否正在处理请求
|
||
|
||
# 异步任务队列(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, np.ndarray], 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 # 队列满回调
|
||
|
||
# 音频播放流
|
||
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[np.ndarray]] = {}
|
||
|
||
async def connect(self):
|
||
"""建立 WebSocket 连接(初始化一次)"""
|
||
if not self.is_connected:
|
||
try:
|
||
self.websocket = await websockets.connect(self.ws_url)
|
||
self.is_connected = True
|
||
print(f"成功连接到 TTS 服务端: {self.ws_url}")
|
||
|
||
# 启动队列消费协程(后台运行)
|
||
asyncio.create_task(self._consume_queue())
|
||
except Exception as e:
|
||
raise ConnectionError(f"连接失败: {str(e)}")
|
||
|
||
async def disconnect(self):
|
||
"""关闭 WebSocket 连接"""
|
||
if self.is_connected and self.websocket:
|
||
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
|
||
|
||
async def add_tts_request(self, tts_text: str, **kwargs) -> str:
|
||
"""
|
||
添加TTS请求到队列(异步非阻塞)
|
||
:param tts_text: 要合成的文本
|
||
:param kwargs: 其他TTS参数
|
||
: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} 已加入队列,当前队列长度: {self.request_queue.qsize()}")
|
||
return req_id
|
||
except asyncio.QueueFull:
|
||
self.on_queue_full(req_id)
|
||
raise Exception(f"队列已满,请求 {req_id} 入队失败")
|
||
|
||
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"开始处理请求 {req_id},剩余队列长度: {self.request_queue.qsize()}")
|
||
|
||
# 处理当前请求
|
||
await self._process_single_request(request)
|
||
|
||
# 标记任务完成
|
||
self.request_queue.task_done()
|
||
self.is_processing = False
|
||
|
||
except Exception as e:
|
||
print(f"队列消费异常: {str(e)}")
|
||
self.is_processing = False
|
||
# 短暂等待后继续消费,避免死循环
|
||
await asyncio.sleep(0.1)
|
||
|
||
async def _process_single_request(self, request: TTSRequest):
|
||
"""处理单个TTS请求"""
|
||
req_id = request.request_id
|
||
|
||
# 参数校验
|
||
if not request.tts_text:
|
||
self.on_error(req_id, "合成文本不能为空")
|
||
return
|
||
|
||
# 确保已连接
|
||
if not self.is_connected:
|
||
try:
|
||
await self.connect()
|
||
except Exception as e:
|
||
self.on_error(req_id, f"连接失败: {str(e)}")
|
||
return
|
||
|
||
# 构造请求数据
|
||
request_data = {
|
||
"tts_text": request.tts_text,
|
||
"mode": request.mode,
|
||
"sft_spk": request.sft_spk,
|
||
"seed": request.seed,
|
||
"stream": request.stream,
|
||
"speed": request.speed,
|
||
"prompt_wav": request.prompt_wav,
|
||
"instruct_text": request.instruct_text
|
||
}
|
||
|
||
try:
|
||
# 发送请求
|
||
await self.websocket.send(json.dumps(request_data))
|
||
|
||
# 处理返回的流式数据
|
||
await self._handle_response(req_id)
|
||
|
||
except websockets.exceptions.ConnectionClosed:
|
||
self.is_connected = False
|
||
self.on_error(req_id, "连接已关闭")
|
||
except Exception as e:
|
||
self.on_error(req_id, f"处理请求失败: {str(e)}")
|
||
|
||
async def _handle_response(self, req_id: str):
|
||
"""处理单个请求的服务端响应"""
|
||
if not self.websocket:
|
||
return
|
||
|
||
try:
|
||
sample_rate = None
|
||
while True:
|
||
# 接收服务端消息(异步)
|
||
response = await self.websocket.recv()
|
||
data = json.loads(response)
|
||
|
||
# 根据状态分发到不同回调(都带req_id)
|
||
status = data.get("status")
|
||
if status == "start":
|
||
# 合成开始 - 返回采样率等信息
|
||
sample_rate = data.get("sample_rate")
|
||
self.on_start(req_id, data)
|
||
|
||
elif status == "stream":
|
||
# 流式音频块 - 转为 numpy 数组
|
||
audio_chunk = np.array(data["audio_chunk"], dtype=np.float32)
|
||
|
||
# 保存到缓冲区
|
||
self.audio_buffers[req_id].append(audio_chunk)
|
||
|
||
# 音频块回调
|
||
self.on_audio_chunk(req_id, audio_chunk)
|
||
|
||
elif status == "end":
|
||
# 合成结束
|
||
self.on_end(req_id, data)
|
||
|
||
# 组装完整音频数据
|
||
full_audio = np.concatenate(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else np.array(
|
||
[], dtype=np.float32)
|
||
|
||
# 准备返回给外部回调的数据
|
||
result_data = {
|
||
"status": "completed",
|
||
"request_id": req_id,
|
||
"sample_rate": sample_rate,
|
||
"audio_data": full_audio,
|
||
"audio_length": len(full_audio),
|
||
"message": data.get("msg", "合成完成")
|
||
}
|
||
|
||
# 调用外部回调(非阻塞)
|
||
if self.external_callback:
|
||
# 使用create_task确保回调不会阻塞主流程
|
||
asyncio.create_task(self._call_external_callback(req_id, result_data))
|
||
|
||
# 清理缓冲区
|
||
if req_id in self.audio_buffers:
|
||
del self.audio_buffers[req_id]
|
||
|
||
break
|
||
|
||
elif status == "error":
|
||
# 错误处理
|
||
error_msg = data.get("msg", "未知错误")
|
||
self.on_error(req_id, error_msg)
|
||
|
||
# 准备错误结果数据
|
||
error_data = {
|
||
"status": "error",
|
||
"request_id": req_id,
|
||
"message": error_msg
|
||
}
|
||
|
||
# 调用外部回调
|
||
if self.external_callback:
|
||
asyncio.create_task(self._call_external_callback(req_id, error_data))
|
||
|
||
# 清理缓冲区
|
||
if req_id in self.audio_buffers:
|
||
del self.audio_buffers[req_id]
|
||
|
||
break
|
||
|
||
except Exception as e:
|
||
error_msg = f"接收数据异常: {str(e)}"
|
||
self.on_error(req_id, error_msg)
|
||
|
||
# 准备错误结果数据
|
||
error_data = {
|
||
"status": "error",
|
||
"request_id": req_id,
|
||
"message": error_msg
|
||
}
|
||
|
||
# 调用外部回调
|
||
if self.external_callback:
|
||
asyncio.create_task(self._call_external_callback(req_id, error_data))
|
||
|
||
# 清理缓冲区
|
||
if req_id in self.audio_buffers:
|
||
del self.audio_buffers[req_id]
|
||
|
||
async def _call_external_callback(self, req_id: str, result_data: Dict[str, Any]):
|
||
"""异步调用外部回调函数"""
|
||
try:
|
||
if self.external_callback:
|
||
# 确保回调是异步的,不会阻塞
|
||
if asyncio.iscoroutinefunction(self.external_callback):
|
||
await self.external_callback(req_id, result_data)
|
||
else:
|
||
# 如果是同步函数,在线程池中执行
|
||
await asyncio.get_event_loop().run_in_executor(
|
||
None, self.external_callback, req_id, result_data
|
||
)
|
||
except Exception as e:
|
||
print(f"外部回调执行异常: {str(e)}")
|
||
|
||
def _play_audio_chunk(self, chunk: np.ndarray):
|
||
"""播放音频块(非阻塞)"""
|
||
if not self.enable_playback or not self.play_stream:
|
||
return
|
||
|
||
try:
|
||
if chunk.size > 0:
|
||
self.play_stream.write(chunk)
|
||
except Exception as e:
|
||
print(f"音频播放异常: {str(e)}")
|
||
|
||
async def wait_queue_empty(self):
|
||
"""等待队列所有任务处理完成(阻塞)"""
|
||
await self.request_queue.join()
|
||
print("所有队列任务已处理完成")
|
||
|
||
|
||
# ------------------------------
|
||
# 使用示例
|
||
# ------------------------------
|
||
class TTSManager:
|
||
"""TTS管理器 - 供外部代码调用"""
|
||
|
||
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts"):
|
||
self.client = CosyVoiceTTSSocketClient(ws_url)
|
||
self._setup_callbacks()
|
||
|
||
def _setup_callbacks(self):
|
||
"""设置内部回调"""
|
||
|
||
# 合成开始回调
|
||
def on_tts_start(req_id, data):
|
||
print(f"\n🎤 开始合成 [{req_id[:8]}] - 采样率: {data.get('sample_rate')}")
|
||
|
||
# 初始化播放流(如果需要播放)
|
||
if self.client.enable_playback:
|
||
sample_rate = data.get("sample_rate", 24000)
|
||
if self.client.play_stream:
|
||
self.client.play_stream.stop()
|
||
self.client.play_stream.close()
|
||
|
||
self.client.play_stream = sd.OutputStream(
|
||
samplerate=sample_rate,
|
||
channels=1,
|
||
dtype=np.float32
|
||
)
|
||
self.client.play_stream.start()
|
||
self.client.current_req_id = req_id
|
||
|
||
# 音频块回调
|
||
def on_audio_chunk(req_id, chunk):
|
||
if req_id == self.client.current_req_id:
|
||
self.client._play_audio_chunk(chunk)
|
||
|
||
# 合成结束回调
|
||
def on_tts_end(req_id, data):
|
||
print(f"🏁 合成完成 [{req_id[:8]}] - {data.get('msg', '完成')}")
|
||
|
||
# 错误回调
|
||
def on_tts_error(req_id, msg):
|
||
print(f"❌ 合成失败 [{req_id[:8]}] - {msg}")
|
||
|
||
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
|
||
|
||
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: 其他参数(mode, sft_spk, speed等)
|
||
:return: 请求ID
|
||
"""
|
||
return await self.client.add_tts_request(text, **kwargs)
|
||
|
||
async def wait_all_completed(self):
|
||
"""等待所有任务完成"""
|
||
await self.client.wait_queue_empty()
|
||
|
||
|
||
# ------------------------------
|
||
# 外部调用示例
|
||
# ------------------------------
|
||
async def external_usage_example():
|
||
"""外部代码使用示例"""
|
||
|
||
# 1. 创建TTS管理器
|
||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||
|
||
# 2. 设置结果回调函数(接收完整结果)
|
||
def handle_tts_result1(req_id: str, result: Dict[str, Any]):
|
||
"""处理TTS结果回调"""
|
||
status = result.get("status")
|
||
|
||
if status == "completed":
|
||
audio_data = result.get("audio_data")
|
||
sample_rate = result.get("sample_rate")
|
||
audio_length = result.get("audio_length")
|
||
|
||
print(f"✅ 收到TTS结果 [{req_id[:8]}]: 长度{audio_length}采样点, 采样率{sample_rate}Hz")
|
||
|
||
# 这里可以保存音频文件或进行其他处理
|
||
# 注意:audio_data是完整的numpy数组
|
||
|
||
elif status == "error":
|
||
error_msg = result.get("message")
|
||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||
|
||
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
||
"""处理TTS结果回调 - 保存为PCM文件"""
|
||
print('处理TTS结果回调', result)
|
||
status = result.get("status")
|
||
print('status')
|
||
if status == "completed":
|
||
audio_data = result.get("audio_data")
|
||
sample_rate = result.get("sample_rate")
|
||
|
||
if audio_data is not None and len(audio_data) > 0:
|
||
# 转换为PCM数据
|
||
pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
|
||
pcm_bytes = pcm_data.tobytes()
|
||
|
||
# 保存为PCM文件
|
||
filename = f"tts_output_{req_id[:8]}.pcm"
|
||
with open(filename, "wb") as f:
|
||
f.write(pcm_bytes)
|
||
|
||
print(f"💾 PCM文件已保存: {filename}, 大小: {len(pcm_bytes)}字节")
|
||
|
||
return filename
|
||
|
||
elif status == "error":
|
||
error_msg = result.get("message")
|
||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||
|
||
return None
|
||
tts_manager.set_result_callback(handle_tts_result)
|
||
|
||
# 3. 设置是否播放(可选,默认True)
|
||
tts_manager.set_playback_enabled(True) # 设置为False则不播放
|
||
|
||
# 4. 初始化连接
|
||
await tts_manager.initialize()
|
||
|
||
# 5. 异步合成多个文本(非阻塞)
|
||
texts = [
|
||
"你好,这是第一个排队的TTS请求。",
|
||
"我是第二个请求,会等第一个处理完再执行。",
|
||
"第三个请求,支持流式播放和队列管理。",
|
||
"第四个请求,测试队列的自动消费功能。",
|
||
"最后一个请求,处理完成后会自动结束。"
|
||
]
|
||
|
||
req_ids = []
|
||
for i, text in enumerate(texts):
|
||
req_id = await tts_manager.synthesize(
|
||
text,
|
||
mode="预训练音色",
|
||
sft_spk="中文女",
|
||
speed=1.0
|
||
)
|
||
req_ids.append(req_id)
|
||
print(f"已提交请求 {i + 1}: ID={req_id[:8]}")
|
||
|
||
# 可以立即继续其他操作,不需要等待
|
||
await asyncio.sleep(0.5) # 模拟其他操作
|
||
|
||
# 6. 可以在这里做其他事情,TTS会在后台处理
|
||
|
||
# 7. 等待所有TTS任务完成(可选)
|
||
await tts_manager.wait_all_completed()
|
||
|
||
# 8. 关闭连接
|
||
await tts_manager.shutdown()
|
||
|
||
|
||
async def advanced_usage_example():
|
||
"""高级使用示例:动态控制播放和结果处理"""
|
||
|
||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||
|
||
# 自定义结果处理器
|
||
class ResultProcessor:
|
||
def __init__(self):
|
||
self.results = {}
|
||
|
||
async def process_result(self, req_id: str, result: Dict[str, Any]):
|
||
"""异步处理结果"""
|
||
if result.get("status") == "completed":
|
||
# 保存音频数据等处理
|
||
self.results[req_id] = result
|
||
print(f"处理完成: {req_id[:8]}, 音频长度: {result.get('audio_length')}")
|
||
|
||
# 这里可以添加自定义逻辑,如保存到文件、发送到网络等
|
||
|
||
processor = ResultProcessor()
|
||
tts_manager.set_result_callback(processor.process_result)
|
||
|
||
# 禁用播放,只获取数据
|
||
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)
|
||
tasks.append(task)
|
||
|
||
# 并行提交所有请求
|
||
req_ids = await asyncio.gather(*tasks)
|
||
|
||
# 等待完成
|
||
await tts_manager.wait_all_completed()
|
||
await tts_manager.shutdown()
|
||
|
||
print(f"所有任务完成,共处理 {len(processor.results)} 个结果")
|
||
|
||
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)}") |