from typing import Optional, List, Dict, Any, Tuple, AsyncGenerator
format='%(asctime)s - %(levelname)s - %(message)s',
logging.FileHandler('run.log'),
logger = logging.getLogger(__name__)
DEFAULT_SAMPLE_RATE = 16000
CLIENT_FULL_REQUEST = 0b0001
CLIENT_AUDIO_ONLY_REQUEST = 0b0010
SERVER_FULL_RESPONSE = 0b1001
SERVER_ERROR_RESPONSE = 0b1111
class MessageTypeSpecificFlags:
NEG_WITH_SEQUENCE = 0b0011
class SerializationType:
NO_SERIALIZATION = 0b0000
# 填入新版控制台获取的 API Key 和 Resource ID
self.api_key = "your_api_key"
self.resource_id = "volc.seedasr.sauc.duration"
def gzip_compress(data: bytes) -> bytes:
return gzip.compress(data)
def gzip_decompress(data: bytes) -> bytes:
return gzip.decompress(data)
def judge_wav(data: bytes) -> bool:
return data[:4] == b'RIFF' and data[8:12] == b'WAVE'
def convert_wav_with_path(audio_path: str, sample_rate: int = DEFAULT_SAMPLE_RATE) -> bytes:
"ffmpeg", "-v", "quiet", "-y", "-i", audio_path,
"-acodec", "pcm_s16le", "-ac", "1", "-ar", str(sample_rate),
result = subprocess.run(cmd, check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
except subprocess.CalledProcessError as e:
logger.error(f"FFmpeg conversion failed: {e.stderr.decode()}")
raise RuntimeError(f"Audio conversion failed: {e.stderr.decode()}")
def read_wav_info(data: bytes) -> Tuple[int, int, int, int, bytes]:
raise ValueError("Invalid WAV file: too short")
raise ValueError("Invalid WAV file: not RIFF format")
raise ValueError("Invalid WAV file: not WAVE format")
audio_format = struct.unpack('<H', data[20:22])[0]
num_channels = struct.unpack('<H', data[22:24])[0]
sample_rate = struct.unpack('<I', data[24:28])[0]
bits_per_sample = struct.unpack('<H', data[34:36])[0]
while pos < len(data) - 8:
subchunk_id = data[pos:pos+4]
subchunk_size = struct.unpack('<I', data[pos+4:pos+8])[0]
if subchunk_id == b'data':
wave_data = data[pos+8:pos+8+subchunk_size]
subchunk_size // (num_channels * (bits_per_sample // 8)),
pos += 8 + subchunk_size
raise ValueError("Invalid WAV file: no data subchunk found")
self.message_type = MessageType.CLIENT_FULL_REQUEST
self.message_type_specific_flags = MessageTypeSpecificFlags.POS_SEQUENCE
self.serialization_type = SerializationType.JSON
self.compression_type = CompressionType.GZIP
self.reserved_data = bytes([0x00])
def with_message_type(self, message_type: int) -> 'AsrRequestHeader':
self.message_type = message_type
def with_message_type_specific_flags(self, flags: int) -> 'AsrRequestHeader':
self.message_type_specific_flags = flags
def with_serialization_type(self, serialization_type: int) -> 'AsrRequestHeader':
self.serialization_type = serialization_type
def with_compression_type(self, compression_type: int) -> 'AsrRequestHeader':
self.compression_type = compression_type
def with_reserved_data(self, reserved_data: bytes) -> 'AsrRequestHeader':
self.reserved_data = reserved_data
def to_bytes(self) -> bytes:
header.append((ProtocolVersion.V1 << 4) | 1)
header.append((self.message_type << 4) | self.message_type_specific_flags)
header.append((self.serialization_type << 4) | self.compression_type)
header.extend(self.reserved_data)
def default_header() -> 'AsrRequestHeader':
return AsrRequestHeader()
def new_auth_headers() -> Dict[str, str]:
reqid = str(uuid.uuid4())
"X-Api-Key": config.api_key,
"X-Api-Resource-Id": config.resource_id,
"X-Api-Request-Id": reqid,
"X-Api-Connect-Id": reqid,
def new_full_client_request(seq: int) -> bytes: # 添加seq参数
header = AsrRequestHeader.default_header() \
.with_message_type_specific_flags(MessageTypeSpecificFlags.POS_SEQUENCE)
"model_name": "bigmodel",
"show_utterances": True,
"enable_nonstream": False
payload_bytes = json.dumps(payload).encode('utf-8')
compressed_payload = CommonUtils.gzip_compress(payload_bytes)
payload_size = len(compressed_payload)
request.extend(header.to_bytes())
request.extend(struct.pack('>i', seq)) # 使用传入的seq
request.extend(struct.pack('>I', payload_size))
request.extend(compressed_payload)
def new_audio_only_request(seq: int, segment: bytes, is_last: bool = False) -> bytes:
header = AsrRequestHeader.default_header()
header.with_message_type_specific_flags(MessageTypeSpecificFlags.NEG_WITH_SEQUENCE)
header.with_message_type_specific_flags(MessageTypeSpecificFlags.POS_SEQUENCE)
header.with_message_type(MessageType.CLIENT_AUDIO_ONLY_REQUEST)
request.extend(header.to_bytes())
request.extend(struct.pack('>i', seq))
compressed_segment = CommonUtils.gzip_compress(segment)
request.extend(struct.pack('>I', len(compressed_segment)))
request.extend(compressed_segment)
self.is_last_package = False
self.payload_sequence = 0
def to_dict(self) -> Dict[str, Any]:
"is_last_package": self.is_last_package,
"payload_sequence": self.payload_sequence,
"payload_size": self.payload_size,
"payload_msg": self.payload_msg
def parse_response(msg: bytes) -> AsrResponse:
response = AsrResponse()
header_size = msg[0] & 0x0f
message_type = msg[1] >> 4
message_type_specific_flags = msg[1] & 0x0f
serialization_method = msg[2] >> 4
message_compression = msg[2] & 0x0f
payload = msg[header_size*4:]
# 解析message_type_specific_flags
if message_type_specific_flags & 0x01:
response.payload_sequence = struct.unpack('>i', payload[:4])[0]
if message_type_specific_flags & 0x02:
response.is_last_package = True
if message_type_specific_flags & 0x04:
response.event = struct.unpack('>i', payload[:4])[0]
if message_type == MessageType.SERVER_FULL_RESPONSE:
response.payload_size = struct.unpack('>I', payload[:4])[0]
elif message_type == MessageType.SERVER_ERROR_RESPONSE:
response.code = struct.unpack('>i', payload[:4])[0]
response.payload_size = struct.unpack('>I', payload[4:8])[0]
if message_compression == CompressionType.GZIP:
payload = CommonUtils.gzip_decompress(payload)
logger.error(f"Failed to decompress payload: {e}")
if serialization_method == SerializationType.JSON:
response.payload_msg = json.loads(payload.decode('utf-8'))
logger.error(f"Failed to parse payload: {e}")
def __init__(self, url: str, segment_duration: int = 200):
self.segment_duration = segment_duration
self.session = None # 添加session引用
async def __aenter__(self):
self.session = aiohttp.ClientSession()
async def __aexit__(self, exc_type, exc, tb):
if self.conn and not self.conn.closed:
if self.session and not self.session.closed:
await self.session.close()
async def read_audio_data(self, file_path: str) -> bytes:
with open(file_path, 'rb') as f:
if not CommonUtils.judge_wav(content):
logger.info("Converting audio to WAV format...")
content = CommonUtils.convert_wav_with_path(file_path, DEFAULT_SAMPLE_RATE)
logger.error(f"Failed to read audio data: {e}")
def get_segment_size(self, content: bytes) -> int:
channel_num, samp_width, frame_rate, _, _ = CommonUtils.read_wav_info(content)[:5]
size_per_sec = channel_num * samp_width * frame_rate
segment_size = size_per_sec * self.segment_duration // 1000
logger.error(f"Failed to calculate segment size: {e}")
async def create_connection(self) -> None:
headers = RequestBuilder.new_auth_headers()
self.conn = await self.session.ws_connect( # 使用self.session
logger.info(f"Connected to {self.url}")
logger.error(f"Failed to connect to WebSocket: {e}")
async def send_full_client_request(self) -> None:
request = RequestBuilder.new_full_client_request(self.seq)
await self.conn.send_bytes(request)
logger.info(f"Sent full client request with seq: {self.seq-1}")
msg = await self.conn.receive()
if msg.type == aiohttp.WSMsgType.BINARY:
response = ResponseParser.parse_response(msg.data)
logger.info(f"Received response: {response.to_dict()}")
logger.error(f"Unexpected message type: {msg.type}")
logger.error(f"Failed to send full client request: {e}")
async def send_messages(self, segment_size: int, content: bytes) -> AsyncGenerator[None, None]:
audio_segments = self.split_audio(content, segment_size)
total_segments = len(audio_segments)
for i, segment in enumerate(audio_segments):
is_last = (i == total_segments - 1)
request = RequestBuilder.new_audio_only_request(
await self.conn.send_bytes(request)
logger.info(f"Sent audio segment with seq: {self.seq} (last: {is_last})")
await asyncio.sleep(self.segment_duration / 1000) # 逐个发送,间隔时间模拟实时流
async def recv_messages(self) -> AsyncGenerator[AsrResponse, None]:
async for msg in self.conn:
if msg.type == aiohttp.WSMsgType.BINARY:
response = ResponseParser.parse_response(msg.data)
if response.is_last_package or response.code != 0:
elif msg.type == aiohttp.WSMsgType.ERROR:
logger.error(f"WebSocket error: {msg.data}")
elif msg.type == aiohttp.WSMsgType.CLOSED:
logger.info("WebSocket connection closed")
logger.error(f"Error receiving messages: {e}")
async def start_audio_stream(self, segment_size: int, content: bytes) -> AsyncGenerator[AsrResponse, None]:
async for _ in self.send_messages(segment_size, content):
sender_task = asyncio.create_task(sender())
async for response in self.recv_messages():
except asyncio.CancelledError:
def split_audio(data: bytes, segment_size: int) -> List[bytes]:
for i in range(0, len(data), segment_size):
segments.append(data[i:end])
async def execute(self, file_path: str) -> AsyncGenerator[AsrResponse, None]:
raise ValueError("File path is empty")
raise ValueError("URL is empty")
content = await self.read_audio_data(file_path)
segment_size = self.get_segment_size(content)
await self.create_connection()
await self.send_full_client_request()
async for response in self.start_audio_stream(segment_size, content):
logger.error(f"Error in ASR execution: {e}")
audio_file_path = "your_file_path" # 替换为你的音频文件路径
ws_url = "wss://openspeech.bytedance.com/api/v3/plan/sauc/bigmodel_nostream"
seg_duration = 200 # 每次发送的音频片段时长(ms)
async with AsrWsClient(ws_url, seg_duration) as client:
async for response in client.execute(audio_file_path):
logger.info(f"Received response: {json.dumps(response.to_dict(), indent=2, ensure_ascii=False)}")
logger.error(f"ASR processing failed: {e}")
if __name__ == "__main__":