2025-06-03 17:30:35 +08:00
|
|
|
|
import time
|
|
|
|
|
|
import os
|
|
|
|
|
|
import uuid
|
2025-02-23 14:38:21 +08:00
|
|
|
|
import json
|
|
|
|
|
|
import gzip
|
2025-05-29 10:38:01 +08:00
|
|
|
|
import websockets
|
2025-05-29 23:56:34 +08:00
|
|
|
|
from config.logger import setup_logging
|
2025-06-03 17:30:35 +08:00
|
|
|
|
from typing import Optional, Tuple, List
|
|
|
|
|
|
from core.providers.asr.base import ASRProviderBase
|
2025-05-29 23:56:34 +08:00
|
|
|
|
from core.providers.asr.dto.dto import InterfaceType
|
2025-06-03 17:30:35 +08:00
|
|
|
|
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
|
|
|
|
|
TAG = __name__
|
|
|
|
|
|
logger = setup_logging()
|
|
|
|
|
|
|
|
|
|
|
|
CLIENT_FULL_REQUEST = 0b0001
|
|
|
|
|
|
CLIENT_AUDIO_ONLY_REQUEST = 0b0010
|
2025-06-03 17:30:35 +08:00
|
|
|
|
|
|
|
|
|
|
NO_SEQUENCE = 0b0000
|
|
|
|
|
|
NEG_SEQUENCE = 0b0010
|
|
|
|
|
|
|
2025-02-23 14:38:21 +08:00
|
|
|
|
SERVER_FULL_RESPONSE = 0b1001
|
|
|
|
|
|
SERVER_ACK = 0b1011
|
|
|
|
|
|
SERVER_ERROR_RESPONSE = 0b1111
|
2025-06-03 17:30:35 +08:00
|
|
|
|
|
|
|
|
|
|
NO_SERIALIZATION = 0b0000
|
|
|
|
|
|
JSON = 0b0001
|
|
|
|
|
|
THRIFT = 0b0011
|
|
|
|
|
|
CUSTOM_TYPE = 0b1111
|
|
|
|
|
|
NO_COMPRESSION = 0b0000
|
|
|
|
|
|
GZIP = 0b0001
|
|
|
|
|
|
CUSTOM_COMPRESSION = 0b1111
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def parse_response(res):
|
|
|
|
|
|
"""
|
|
|
|
|
|
protocol_version(4 bits), header_size(4 bits),
|
|
|
|
|
|
message_type(4 bits), message_type_specific_flags(4 bits)
|
|
|
|
|
|
serialization_method(4 bits) message_compression(4 bits)
|
|
|
|
|
|
reserved (8bits) 保留字段
|
|
|
|
|
|
header_extensions 扩展头(大小等于 8 * 4 * (header_size - 1) )
|
|
|
|
|
|
payload 类似与http 请求体
|
|
|
|
|
|
"""
|
|
|
|
|
|
protocol_version = res[0] >> 4
|
|
|
|
|
|
header_size = res[0] & 0x0F
|
|
|
|
|
|
message_type = res[1] >> 4
|
|
|
|
|
|
message_type_specific_flags = res[1] & 0x0F
|
|
|
|
|
|
serialization_method = res[2] >> 4
|
|
|
|
|
|
message_compression = res[2] & 0x0F
|
|
|
|
|
|
reserved = res[3]
|
|
|
|
|
|
header_extensions = res[4 : header_size * 4]
|
|
|
|
|
|
payload = res[header_size * 4 :]
|
|
|
|
|
|
result = {}
|
|
|
|
|
|
payload_msg = None
|
|
|
|
|
|
payload_size = 0
|
|
|
|
|
|
if message_type == SERVER_FULL_RESPONSE:
|
|
|
|
|
|
payload_size = int.from_bytes(payload[:4], "big", signed=True)
|
|
|
|
|
|
payload_msg = payload[4:]
|
|
|
|
|
|
elif message_type == SERVER_ACK:
|
|
|
|
|
|
seq = int.from_bytes(payload[:4], "big", signed=True)
|
|
|
|
|
|
result["seq"] = seq
|
|
|
|
|
|
if len(payload) >= 8:
|
|
|
|
|
|
payload_size = int.from_bytes(payload[4:8], "big", signed=False)
|
|
|
|
|
|
payload_msg = payload[8:]
|
|
|
|
|
|
elif message_type == SERVER_ERROR_RESPONSE:
|
|
|
|
|
|
code = int.from_bytes(payload[:4], "big", signed=False)
|
|
|
|
|
|
result["code"] = code
|
|
|
|
|
|
payload_size = int.from_bytes(payload[4:8], "big", signed=False)
|
|
|
|
|
|
payload_msg = payload[8:]
|
|
|
|
|
|
if payload_msg is None:
|
|
|
|
|
|
return result
|
|
|
|
|
|
if message_compression == GZIP:
|
|
|
|
|
|
payload_msg = gzip.decompress(payload_msg)
|
|
|
|
|
|
if serialization_method == JSON:
|
|
|
|
|
|
payload_msg = json.loads(str(payload_msg, "utf-8"))
|
|
|
|
|
|
elif serialization_method != NO_SERIALIZATION:
|
|
|
|
|
|
payload_msg = str(payload_msg, "utf-8")
|
|
|
|
|
|
result["payload_msg"] = payload_msg
|
|
|
|
|
|
result["payload_size"] = payload_size
|
|
|
|
|
|
return result
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class ASRProvider(ASRProviderBase):
|
2025-06-03 17:30:35 +08:00
|
|
|
|
def __init__(self, config: dict, delete_audio_file: bool):
|
2025-05-08 15:08:11 +08:00
|
|
|
|
super().__init__()
|
2025-06-03 17:30:35 +08:00
|
|
|
|
self.interface_type = InterfaceType.NON_STREAM
|
|
|
|
|
|
self.appid = config.get("appid")
|
2025-02-23 14:38:21 +08:00
|
|
|
|
self.cluster = config.get("cluster")
|
|
|
|
|
|
self.access_token = config.get("access_token")
|
2025-05-08 12:07:50 +08:00
|
|
|
|
self.boosting_table_name = config.get("boosting_table_name", "")
|
|
|
|
|
|
self.correct_table_name = config.get("correct_table_name", "")
|
2025-06-03 17:30:35 +08:00
|
|
|
|
self.output_dir = config.get("output_dir")
|
2025-04-29 01:28:06 +08:00
|
|
|
|
self.delete_audio_file = delete_audio_file
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
self.host = "openspeech.bytedance.com"
|
|
|
|
|
|
self.ws_url = f"wss://{self.host}/api/v2/asr"
|
|
|
|
|
|
self.success_code = 1000
|
|
|
|
|
|
self.seg_duration = 15000
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
# 确保输出目录存在
|
|
|
|
|
|
os.makedirs(self.output_dir, exist_ok=True)
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
@staticmethod
|
|
|
|
|
|
def _generate_header(
|
|
|
|
|
|
message_type=CLIENT_FULL_REQUEST, message_type_specific_flags=NO_SEQUENCE
|
|
|
|
|
|
) -> bytearray:
|
|
|
|
|
|
"""Generate protocol header."""
|
|
|
|
|
|
header = bytearray()
|
|
|
|
|
|
header_size = 1
|
|
|
|
|
|
header.append((0b0001 << 4) | header_size) # Protocol version
|
|
|
|
|
|
header.append((message_type << 4) | message_type_specific_flags)
|
|
|
|
|
|
header.append((0b0001 << 4) | 0b0001) # JSON serialization & GZIP compression
|
|
|
|
|
|
header.append(0x00) # reserved
|
|
|
|
|
|
return header
|
2025-05-30 15:47:32 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
def _construct_request(self, reqid) -> dict:
|
|
|
|
|
|
"""Construct the request payload."""
|
|
|
|
|
|
return {
|
2025-02-23 14:38:21 +08:00
|
|
|
|
"app": {
|
2025-06-03 17:30:35 +08:00
|
|
|
|
"appid": f"{self.appid}",
|
2025-02-23 14:38:21 +08:00
|
|
|
|
"cluster": self.cluster,
|
|
|
|
|
|
"token": self.access_token,
|
|
|
|
|
|
},
|
2025-06-03 17:30:35 +08:00
|
|
|
|
"user": {
|
|
|
|
|
|
"uid": str(uuid.uuid4()),
|
|
|
|
|
|
},
|
2025-04-30 13:34:00 +08:00
|
|
|
|
"request": {
|
2025-05-03 22:20:50 +08:00
|
|
|
|
"reqid": reqid,
|
2025-06-03 17:30:35 +08:00
|
|
|
|
"show_utterances": False,
|
2025-04-30 13:34:00 +08:00
|
|
|
|
"sequence": 1,
|
|
|
|
|
|
"boosting_table_name": self.boosting_table_name,
|
|
|
|
|
|
"correct_table_name": self.correct_table_name,
|
|
|
|
|
|
},
|
2025-02-23 14:38:21 +08:00
|
|
|
|
"audio": {
|
2025-06-03 17:30:35 +08:00
|
|
|
|
"format": "raw",
|
|
|
|
|
|
"rate": 16000,
|
|
|
|
|
|
"language": "zh-CN",
|
|
|
|
|
|
"bits": 16,
|
|
|
|
|
|
"channel": 1,
|
|
|
|
|
|
"codec": "raw",
|
2025-02-23 14:38:21 +08:00
|
|
|
|
},
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
async def _send_request(
|
|
|
|
|
|
self, audio_data: List[bytes], segment_size: int
|
|
|
|
|
|
) -> Optional[str]:
|
|
|
|
|
|
"""Send request to Volcano ASR service."""
|
2025-05-30 15:47:32 +08:00
|
|
|
|
try:
|
2025-06-03 17:30:35 +08:00
|
|
|
|
auth_header = {"Authorization": "Bearer; {}".format(self.access_token)}
|
|
|
|
|
|
async with websockets.connect(
|
|
|
|
|
|
self.ws_url, additional_headers=auth_header
|
|
|
|
|
|
) as websocket:
|
|
|
|
|
|
# Prepare request data
|
|
|
|
|
|
request_params = self._construct_request(str(uuid.uuid4()))
|
|
|
|
|
|
payload_bytes = str.encode(json.dumps(request_params))
|
|
|
|
|
|
payload_bytes = gzip.compress(payload_bytes)
|
|
|
|
|
|
full_client_request = self._generate_header()
|
|
|
|
|
|
full_client_request.extend(
|
|
|
|
|
|
(len(payload_bytes)).to_bytes(4, "big")
|
|
|
|
|
|
) # payload size(4 bytes)
|
|
|
|
|
|
full_client_request.extend(payload_bytes) # payload
|
2025-05-30 15:47:32 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
# Send header and metadata
|
|
|
|
|
|
# full_client_request
|
|
|
|
|
|
await websocket.send(full_client_request)
|
|
|
|
|
|
res = await websocket.recv()
|
|
|
|
|
|
result = parse_response(res)
|
|
|
|
|
|
if (
|
|
|
|
|
|
"payload_msg" in result
|
|
|
|
|
|
and result["payload_msg"]["code"] != self.success_code
|
2025-06-08 01:09:32 +08:00
|
|
|
|
and result["payload_msg"]["code"] != 1013 # 忽略无有效语音的错误
|
2025-06-03 17:30:35 +08:00
|
|
|
|
):
|
|
|
|
|
|
logger.bind(tag=TAG).error(f"ASR error: {result}")
|
|
|
|
|
|
return None
|
2025-05-29 23:56:34 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
for seq, (chunk, last) in enumerate(
|
|
|
|
|
|
self.slice_data(audio_data, segment_size), 1
|
|
|
|
|
|
):
|
|
|
|
|
|
if last:
|
|
|
|
|
|
audio_only_request = self._generate_header(
|
|
|
|
|
|
message_type=CLIENT_AUDIO_ONLY_REQUEST,
|
|
|
|
|
|
message_type_specific_flags=NEG_SEQUENCE,
|
2025-02-23 14:38:21 +08:00
|
|
|
|
)
|
2025-06-03 17:30:35 +08:00
|
|
|
|
else:
|
|
|
|
|
|
audio_only_request = self._generate_header(
|
|
|
|
|
|
message_type=CLIENT_AUDIO_ONLY_REQUEST
|
|
|
|
|
|
)
|
|
|
|
|
|
payload_bytes = gzip.compress(chunk)
|
|
|
|
|
|
audio_only_request.extend(
|
|
|
|
|
|
(len(payload_bytes)).to_bytes(4, "big")
|
|
|
|
|
|
) # payload size(4 bytes)
|
|
|
|
|
|
audio_only_request.extend(payload_bytes) # payload
|
|
|
|
|
|
# Send audio data
|
|
|
|
|
|
await websocket.send(audio_only_request)
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
# Receive response
|
|
|
|
|
|
response = await websocket.recv()
|
|
|
|
|
|
result = parse_response(response)
|
2025-05-30 15:47:32 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
if (
|
|
|
|
|
|
"payload_msg" in result
|
|
|
|
|
|
and result["payload_msg"]["code"] == self.success_code
|
|
|
|
|
|
):
|
|
|
|
|
|
if len(result["payload_msg"]["result"]) > 0:
|
|
|
|
|
|
return result["payload_msg"]["result"][0]["text"]
|
|
|
|
|
|
return None
|
2025-06-08 01:09:32 +08:00
|
|
|
|
elif "payload_msg" in result and result["payload_msg"]["code"] == 1013:
|
|
|
|
|
|
# 无有效语音,返回空字符串
|
|
|
|
|
|
return ""
|
2025-06-03 17:30:35 +08:00
|
|
|
|
else:
|
|
|
|
|
|
logger.bind(tag=TAG).error(f"ASR error: {result}")
|
|
|
|
|
|
return None
|
2025-05-30 15:47:32 +08:00
|
|
|
|
|
2025-02-23 14:38:21 +08:00
|
|
|
|
except Exception as e:
|
2025-06-03 17:30:35 +08:00
|
|
|
|
logger.bind(tag=TAG).error(f"ASR request failed: {e}", exc_info=True)
|
|
|
|
|
|
return None
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
@staticmethod
|
|
|
|
|
|
def slice_data(data: bytes, chunk_size: int) -> (list, bool):
|
2025-02-23 14:38:21 +08:00
|
|
|
|
"""
|
2025-06-03 17:30:35 +08:00
|
|
|
|
slice data
|
|
|
|
|
|
:param data: wav data
|
|
|
|
|
|
:param chunk_size: the segment size in one request
|
|
|
|
|
|
:return: segment data, last flag
|
2025-02-23 14:38:21 +08:00
|
|
|
|
"""
|
2025-06-03 17:30:35 +08:00
|
|
|
|
data_len = len(data)
|
|
|
|
|
|
offset = 0
|
|
|
|
|
|
while offset + chunk_size < data_len:
|
|
|
|
|
|
yield data[offset : offset + chunk_size], False
|
|
|
|
|
|
offset += chunk_size
|
2025-02-23 14:38:21 +08:00
|
|
|
|
else:
|
2025-06-03 17:30:35 +08:00
|
|
|
|
yield data[offset:data_len], True
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
async def speech_to_text(
|
2026-01-29 16:39:35 +08:00
|
|
|
|
self, opus_data: List[bytes], session_id: str, audio_format="opus", artifacts=None
|
2025-06-03 17:30:35 +08:00
|
|
|
|
) -> Tuple[Optional[str], Optional[str]]:
|
|
|
|
|
|
"""将语音数据转换为文本"""
|
2025-05-03 22:20:50 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
try:
|
2026-01-25 11:27:34 +08:00
|
|
|
|
if artifacts is None:
|
|
|
|
|
|
return "", None
|
2025-04-29 01:28:06 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
# 直接使用PCM数据
|
|
|
|
|
|
# 计算分段大小 (单声道, 16bit, 16kHz采样率)
|
|
|
|
|
|
size_per_sec = 1 * 2 * 16000 # nchannels * sampwidth * framerate
|
|
|
|
|
|
segment_size = int(size_per_sec * self.seg_duration / 1000)
|
2025-02-23 14:38:21 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
# 语音识别
|
|
|
|
|
|
start_time = time.time()
|
2026-01-25 11:27:34 +08:00
|
|
|
|
text = await self._send_request(artifacts.pcm_bytes, segment_size)
|
2025-06-03 17:30:35 +08:00
|
|
|
|
if text:
|
|
|
|
|
|
logger.bind(tag=TAG).debug(
|
|
|
|
|
|
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}"
|
2025-05-30 15:47:32 +08:00
|
|
|
|
)
|
2026-01-25 11:27:34 +08:00
|
|
|
|
return text, artifacts.file_path
|
|
|
|
|
|
return "", artifacts.file_path
|
2025-05-30 15:47:32 +08:00
|
|
|
|
|
2025-06-03 17:30:35 +08:00
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
|
2026-01-25 11:27:34 +08:00
|
|
|
|
return "", None
|