test/scripts/benchmark/clients/tts_client.py

252 lines
8.5 KiB
Python

# -*- coding: utf-8 -*-
"""
TTS WebSocket 测试客户端
参考 realtime-llm/backend/cti_websocket_handler.py 中的调用方式。
模拟 LLM 流式输出场景,按句子发送文本进行合成。
"""
import asyncio
import json
import time
import logging
import wave
from pathlib import Path
from typing import Optional
from .base_client import BaseWebSocketClient
from ..metrics.models import TTSMetrics
logger = logging.getLogger(__name__)
# 协议常量
TTS_NAMESPACE = "FlowingSpeechSynthesizer"
MSG_START_SYNTHESIS = "StartSynthesis"
MSG_RUN_SYNTHESIS = "RunSynthesis"
MSG_STOP_SYNTHESIS = "StopSynthesis"
MSG_SYNTHESIS_STARTED = "SynthesisStarted"
MSG_SENTENCE_BEGIN = "SentenceBegin"
MSG_SENTENCE_END = "SentenceEnd"
MSG_SYNTHESIS_COMPLETED = "SynthesisCompleted"
MSG_TASK_FAILED = "TaskFailed"
class TTSWebSocketClient(BaseWebSocketClient):
"""TTS WebSocket 测试客户端 (模拟流式文本输入)"""
def __init__(
self,
ws_url: str,
text: str,
voice: str = "中文女",
audio_format: str = "PCM",
sample_rate: int = 22050,
timeout: float = 120.0,
chunk_interval: float = 0.05, # 发送间隔 (秒),模拟 LLM 生成速度
debug: bool = False, # 调试模式
save_audio_dir: Optional[Path] = None, # 保存音频的目录
):
super().__init__(ws_url, timeout)
self.text = text
self.voice = voice
self.audio_format = audio_format
self.sample_rate = sample_rate
self.chunk_interval = chunk_interval
self.debug = debug
self.save_audio_dir = save_audio_dir
self._audio_chunks = [] # 存储接收到的音频数据
def _log(self, msg: str):
"""调试日志"""
if self.debug:
logger.info(f"[{self.task_id[:8]}] {msg}")
async def run_test(self) -> TTSMetrics:
"""执行 TTS 测试"""
metrics = TTSMetrics(
request_id=self.task_id,
concurrency_level=0,
start_time=time.perf_counter(),
text_length=len(self.text),
sample_rate=self.sample_rate,
)
try:
await asyncio.wait_for(
self._run_tts_session(metrics),
timeout=self.timeout,
)
metrics.success = True
except asyncio.TimeoutError:
metrics.error_message = "Timeout"
logger.warning(f"TTS 请求超时: {self.task_id}")
except Exception as e:
metrics.error_message = str(e)
logger.warning(f"TTS 请求失败: {self.task_id}, 错误: {e}")
finally:
await self.close()
return metrics
async def _run_tts_session(self, metrics: TTSMetrics) -> None:
"""运行完整的 TTS 会话"""
self._log(f"连接 {self.ws_url}")
await self.connect()
# 用于同步的事件
started_event = asyncio.Event()
completed_event = asyncio.Event()
error_message = None
# 1. 发送 StartSynthesis
await self._send_start_synthesis()
# 2. 启动接收任务
async def receive_loop():
nonlocal error_message
while True:
try:
response = await self.receive()
except Exception as e:
self._log(f"接收异常: {e}")
break
if isinstance(response, bytes):
if metrics.first_chunk_time is None:
metrics.first_chunk_time = time.perf_counter()
metrics.audio_bytes_received += len(response)
# 收集音频数据用于保存
if self.save_audio_dir:
self._audio_chunks.append(response)
self._log(f"← 收到音频: {len(response)} bytes")
elif isinstance(response, str):
try:
data = json.loads(response)
header = data.get("header", {})
name = header.get("name", "")
status = header.get("status", 0)
self._log(f"← 收到事件: {name} (status={status})")
if name == MSG_SYNTHESIS_STARTED:
started_event.set()
elif name == MSG_SENTENCE_END:
metrics.sentence_end_time = time.perf_counter()
elif name == MSG_SYNTHESIS_COMPLETED:
completed_event.set()
break
elif name == MSG_TASK_FAILED:
status_text = header.get("status_text", "Unknown error")
error_message = f"TaskFailed: {status_text}"
self._log(f"← 错误: {status_text}")
completed_event.set()
break
except json.JSONDecodeError:
pass
receive_task = asyncio.create_task(receive_loop())
try:
# 3. 等待 SynthesisStarted
self._log("等待 SynthesisStarted...")
await asyncio.wait_for(started_event.wait(), timeout=10.0)
self._log("收到 SynthesisStarted")
# 4. 发送文本 - 直接发送完整文本,不分割
# 参考 CTI 客户端:发送完整句子而不是切分片段
await self._send_run_synthesis(self.text)
# 5. 发送 StopSynthesis
await self._send_stop_synthesis()
# 6. 等待 SynthesisCompleted
self._log("等待 SynthesisCompleted...")
await completed_event.wait()
if error_message:
raise Exception(error_message)
metrics.complete_time = time.perf_counter()
self._log(f"完成! 收到 {metrics.audio_bytes_received} bytes 音频")
# 保存音频文件
if self.save_audio_dir and self._audio_chunks:
self._save_audio()
finally:
if not receive_task.done():
receive_task.cancel()
try:
await receive_task
except asyncio.CancelledError:
pass
async def _send_start_synthesis(self) -> None:
"""发送 StartSynthesis 消息"""
message = {
"header": self._create_header(MSG_START_SYNTHESIS, TTS_NAMESPACE),
"payload": {
"voice": self.voice,
"format": self.audio_format,
"sample_rate": self.sample_rate,
"volume": 50,
"speech_rate": 0,
"pitch_rate": 0,
"platform": "python",
},
}
self._log(f"→ 发送 StartSynthesis (voice={self.voice}, format={self.audio_format})")
await self.send_json(message)
async def _send_run_synthesis(self, text: str) -> None:
"""发送 RunSynthesis 消息"""
message = {
"header": self._create_header(MSG_RUN_SYNTHESIS, TTS_NAMESPACE),
"payload": {
"text": text,
},
}
# 截断显示
display_text = text[:50] + "..." if len(text) > 50 else text
self._log(f"→ 发送 RunSynthesis: \"{display_text}\" ({len(text)} chars)")
await self.send_json(message)
async def _send_stop_synthesis(self) -> None:
"""发送 StopSynthesis 消息"""
message = {
"header": self._create_header(MSG_STOP_SYNTHESIS, TTS_NAMESPACE),
}
self._log("→ 发送 StopSynthesis")
await self.send_json(message)
def _save_audio(self) -> None:
"""保存收到的音频数据为 WAV 文件"""
if self.save_audio_dir is None:
return
try:
# 合并所有音频块
audio_data = b"".join(self._audio_chunks)
if not audio_data:
return
# 生成文件名
filename = f"{self.task_id[:8]}_{len(self.text)}chars.wav"
filepath = self.save_audio_dir / filename
# PCM 数据保存为 WAV
with wave.open(str(filepath), 'wb') as wav_file:
wav_file.setnchannels(1) # 单声道
wav_file.setsampwidth(2) # 16位 = 2字节
wav_file.setframerate(self.sample_rate)
wav_file.writeframes(audio_data)
self._log(f"音频已保存: {filepath}")
except Exception as e:
logger.warning(f"保存音频失败: {e}")