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)}")