82 lines
2.4 KiB
Python
82 lines
2.4 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""
|
|
测试配置模块
|
|
"""
|
|
|
|
from dataclasses import dataclass, field
|
|
from typing import List, Optional
|
|
from pathlib import Path
|
|
|
|
|
|
@dataclass
|
|
class TestConfig:
|
|
"""测试配置"""
|
|
|
|
# 服务器配置
|
|
host: str = "localhost"
|
|
port: int = 8000
|
|
timeout_seconds: float = 300.0 # 默认 5 分钟,并发 TTS 可能需要更长时间
|
|
warmup_requests: int = 3
|
|
|
|
# 并发配置
|
|
concurrency_levels: List[int] = field(default_factory=lambda: [5, 10, 20, 50])
|
|
|
|
# ASR 配置
|
|
asr_audio_file: Optional[Path] = None
|
|
asr_sample_rate: int = 16000
|
|
asr_chunk_size: int = 9600 # 600ms @ 16kHz
|
|
asr_format: str = "pcm"
|
|
|
|
# TTS 配置
|
|
tts_text_count: int = 50 # 预生成的测试文本数量
|
|
tts_text_length_range: tuple = (50, 100) # 文本字符数范围
|
|
tts_voice: str = "中文女"
|
|
tts_format: str = "PCM"
|
|
tts_sample_rate: int = 22050
|
|
tts_chunk_interval: float = 0.05 # 发送间隔秒数 (模拟 LLM 生成速度)
|
|
|
|
# 输出配置
|
|
output_dir: Path = field(default_factory=lambda: Path("./benchmark_results"))
|
|
report_name: str = "benchmark_report"
|
|
|
|
@property
|
|
def ws_base_url(self) -> str:
|
|
"""WebSocket 基础 URL"""
|
|
return f"ws://{self.host}:{self.port}"
|
|
|
|
@property
|
|
def asr_ws_url(self) -> str:
|
|
"""ASR WebSocket URL"""
|
|
return f"{self.ws_base_url}/ws/v1/asr"
|
|
|
|
@property
|
|
def tts_ws_url(self) -> str:
|
|
"""TTS WebSocket URL"""
|
|
return f"{self.ws_base_url}/ws/v1/tts"
|
|
|
|
def validate(self, test_type: str = "both") -> None:
|
|
"""
|
|
验证配置
|
|
|
|
Args:
|
|
test_type: 测试类型 (asr/tts/both)
|
|
|
|
Raises:
|
|
ValueError: 配置无效
|
|
"""
|
|
if test_type in ("asr", "both"):
|
|
if self.asr_audio_file is None:
|
|
raise ValueError("ASR 测试需要提供音频文件路径 (--audio-file)")
|
|
if not self.asr_audio_file.exists():
|
|
raise ValueError(f"音频文件不存在: {self.asr_audio_file}")
|
|
|
|
if not self.concurrency_levels:
|
|
raise ValueError("至少需要一个并发级别")
|
|
|
|
for level in self.concurrency_levels:
|
|
if level < 1:
|
|
raise ValueError(f"并发级别必须大于 0: {level}")
|
|
|
|
if self.timeout_seconds <= 0:
|
|
raise ValueError("超时时间必须大于 0")
|