Files
banban/talkingq-url/test/aliyun_asr_simple.py
2026-04-16 21:13:06 +08:00

538 lines
19 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import asyncio
from test.audio_processor import AudioProcessor
import math
import asyncio
import struct
import websockets
from dashscope.audio.asr import TranslationRecognizerChat
from dashscope.audio.asr import TranscriptionResult
from dashscope.audio.asr import TranslationRecognizerCallback
# from config import settings
from providers.asr import ASR
# from providers.impl.base_service import BaseService
from providers.impl.interrupt_handler import interrupt_handler
from utils.logger import session_logger
class GummyCallback(TranslationRecognizerCallback):
def __init__(self, device_id, session_id):
self.device_id = device_id
self.session_id = session_id
self.transcript = ""
self.transcript_event = asyncio.Event()
self.received_result = False
self.error_message = None
self.is_closed = False
self.session_end = False
def on_open(self) -> None:
session_logger.info(self.device_id, self.session_id, "Gummy ASR连接建立成功")
def on_event(
self,
request_id,
transcription_result: TranscriptionResult,
translation_result,
usage,
) -> None:
if transcription_result:
if transcription_result.is_sentence_end:
self.transcript = transcription_result.text
self.received_result = True
session_logger.info(
self.device_id,
self.session_id,
f"Gummy ASR完整句子: {self.transcript}",
)
self.transcript_event.set()
else:
self.transcript = transcription_result.text
def on_complete(self) -> None:
session_logger.info(self.device_id, self.session_id, "Gummy ASR识别完成")
self.session_end = True
self.transcript_event.set()
def on_error(self, result) -> None:
self.error_message = str(result)
session_logger.error(
self.device_id, self.session_id, f"Gummy ASR错误: {self.error_message}"
)
self.transcript_event.set()
def on_close(self) -> None:
session_logger.info(self.device_id, self.session_id, "Gummy ASR连接已关闭")
self.is_closed = True
self.transcript_event.set()
class AliyunASR(ASR):
def __init__(self, selected_role=None):
# BaseService.__init__(self, "asr", selected_role)
self.api_key = 'sk-7a50eca6856d4afb968ac3bf512f6d1b'
self.vocabulary_id = 'vocab-talkingq-8c1ed0b3b76342819a1d39988ff22a17'
self.connected = False
self.transcript = ""
self.device_id = "unknown"
self.session_id = "unknown"
self.sample_rate = 16000
self.audio_buffer = bytearray()
self.translator = None
self.callback = None
self.buffer_size = 3200 # 约100ms的音频按16000采样率计算
self.max_end_silence = 2000 # 设为2000ms允许更长的停顿
async def connect(self):
try:
if self.connected and self.translator:
session_logger.info(
self.device_id, self.session_id, "Gummy ASR连接已存在复用当前连接"
)
return
self.transcript = ""
self.audio_buffer = bytearray()
self.callback = GummyCallback(self.device_id, self.session_id)
self.translator = TranslationRecognizerChat(
model="gummy-chat-v1",
format="mp3",
sample_rate=self.sample_rate,
callback=self.callback,
transcription_enabled=True,
source_language="auto",
semantic_punctuation_enabled=False,
max_end_silence=self.max_end_silence,
api_key=self.api_key,
vocabulary_id=self.vocabulary_id,
)
self.translator.start()
self.callback.transcript_event.clear()
self.connected = True
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"建立Gummy ASR连接失败: {str(e)}"
)
self.connected = False
if self.translator:
try:
self.translator.stop()
except:
pass
self.translator = None
async def send_audio(self, audio_data: bytes):
if not self.connected or not self.translator:
session_logger.error(self.device_id, self.session_id, "Gummy ASR连接未建立")
return
try:
self.audio_buffer.extend(audio_data)
while len(self.audio_buffer) >= self.buffer_size:
chunk = bytes(self.audio_buffer[: self.buffer_size])
self.audio_buffer = self.audio_buffer[self.buffer_size :]
if self.device_id and self.session_id:
session_key = (self.device_id, self.session_id)
if interrupt_handler.is_interrupted(session_key):
session_logger.info(
self.device_id,
self.session_id,
"检测到中断,停止发送音频数据",
)
return
try:
self.translator.send_audio_frame(chunk)
except Exception as e:
session_logger.error(
self.device_id,
self.session_id,
f"发送音频帧到Gummy ASR出错: {str(e)}",
)
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"发送音频到Gummy ASR出错: {str(e)}"
)
async def send_end(self):
if not self.connected or not self.translator:
session_logger.error(
self.device_id, self.session_id, "Gummy ASR连接未建立无法发送结束信号"
)
return
try:
if len(self.audio_buffer) > 0:
chunk = bytes(self.audio_buffer)
self.audio_buffer = bytearray()
self.translator.send_audio_frame(chunk)
session_logger.info(
self.device_id,
self.session_id,
"客户端发送结束包等待Gummy ASR处理完毕",
)
loop = asyncio.get_running_loop()
await loop.run_in_executor(None, lambda: self._safe_stop_translator())
except Exception as e:
session_logger.error(
self.device_id,
self.session_id,
f"发送结束信号到Gummy ASR出错: {str(e)}",
)
async def receive_results(self):
if not self.connected:
session_logger.error(
self.device_id, self.session_id, "Gummy ASR连接未建立无法接收结果"
)
return
try:
try:
await asyncio.wait_for(
self.callback.transcript_event.wait(), timeout=0.5
)
except asyncio.TimeoutError:
session_logger.warning(
self.device_id,
self.session_id,
"接收Gummy ASR结果超时使用最后的中间结果",
)
if self.callback.received_result or self.callback.transcript:
self.transcript = self.callback.transcript
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"接收Gummy ASR结果出错: {str(e)}"
)
return self.transcript
async def close(self):
if self.connected:
try:
if self.translator is not None:
try:
self.translator.stop()
except Exception as e:
if "has stopped" not in str(e):
session_logger.warning(
self.device_id,
self.session_id,
f"关闭Gummy translator出错: {str(e)}"
)
self.translator = None
self.connected = False
session_logger.info(
self.device_id, self.session_id, "Gummy ASR连接已关闭"
)
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"关闭Gummy ASR连接出错: {str(e)}"
)
def _safe_stop_translator(self):
"""安全停止translator忽略已停止的异常"""
try:
if self.translator:
self.translator.stop()
except Exception as e:
if "has stopped" not in str(e):
session_logger.warning(
self.device_id,
self.session_id,
f"停止translator时出现非标准错误: {str(e)}",
)
return None
async def transcribe_stream(asr: AliyunASR, stream_generator):
"""
流式转录音频数据
stream_generator: 生成器函数,每次产生音频块(bytes)
返回: 完整转录文本
"""
try:
# 连接ASR服务
await asr.connect()
if not asr.connected:
raise ConnectionError("无法连接到ASR服务")
# 处理音频流
for audio_chunk in stream_generator:
# 检查中断
# if interrupt_handler.is_interrupted((asr.device_id, asr.session_id)):
# session_logger.info(
# asr.device_id,
# asr.session_id,
# "检测到中断,停止音频发送"
# )
# return ""
await asr.send_audio(audio_chunk)
# 可选:获取中间结果
intermediate = await asr.receive_results()
if intermediate:
yield intermediate
# 发送结束信号并获取最终结果
await asr.send_end()
# return await asr.receive_results()
except Exception as e:
# session_logger.error(
# asr.device_id,
# asr.session_id,
# f"流式转录出错: {str(e)}"
# )
print(e)
yield ""
finally:
await asr.close()
# def audio_chunk_generator1(file_path, chunk_duration=100):
# """
# 生成器函数:从音频文件产生指定时长的块
# chunk_duration - 毫秒
# """
# audio = AudioProcessor.load_audio(file_path)
# chunk_size = int(16000 * 2 * (chunk_duration / 1000)) # 16kHz * 2字节 * 时长
# # 获取PCM数据
# pcm_data = AudioProcessor.audio_to_pcm(audio)
# # 分块产生
# for i in range(0, len(pcm_data), chunk_size):
# yield pcm_data[i:i+chunk_size]
# # 模拟实时延迟
# # await asyncio.sleep(chunk_duration / 1000)
def audio_chunk_generator(file_path, chunk_duration=100):
"""
生成器函数:从音频文件产生指定时长的块
chunk_duration - 毫秒
特点直接从AudioSegment对象生成音频块无需额外PCM转换
"""
# 加载音频文件
# audio = AudioProcessor.load_audio(file_path)
with open(file_path, "rb") as f:
audio = f.read()
# 计算每秒的字节数 (sample_rate * sample_width * channels)
# bytes_per_second = audio.frame_rate * (audio.sample_width) * audio.channels
# 计算每个块的字节数 (按时间计算而非固定值)
chunk_size_bytes = 320 #int(bytes_per_second * (chunk_duration / 1000.0))
# 计算块的数量
total_bytes = len(audio)
num_chunks = math.ceil(total_bytes / chunk_size_bytes)
# 分块产生音频数据
for i in range(num_chunks):
# 计算当前块的起始和结束位置
start_byte = i * chunk_size_bytes
end_byte = min((i + 1) * chunk_size_bytes, total_bytes)
# 提取原始字节数据
chunk_data = audio[start_byte:end_byte]
yield chunk_data
def create_websocket_packet(device_id: str, session_id: str, sequence_number: int, packet_type: int, audio_data: bytes) -> bytes:
"""
创建符合WebSocket接口格式的数据包
"""
# 设备ID以null结尾
device_id_bytes = device_id.encode('ascii') + b'\x00'
# 会话ID固定33字节不足补null
session_id_bytes = session_id.encode('ascii')
session_id_bytes = session_id_bytes.ljust(33, b'\x00')[:33]
# 序列号4字节小端
seq_bytes = struct.pack('<I', sequence_number)
# 包类型1字节
type_bytes = struct.pack('<B', packet_type)
# 音频数据大小4字节小端
size_bytes = struct.pack('<I', len(audio_data))
# 组合数据包
packet = device_id_bytes + session_id_bytes + seq_bytes + type_bytes + size_bytes + audio_data
return packet
async def send_audio_to_websocket(file_path: str, websocket_url: str = "ws://localhost:8000/ws"):
"""
将音频文件通过WebSocket发送到指定接口并接收打印返回结果
"""
device_id = "TalkingQ_10003B58862C"
session_id = "test-session-001"
# 加载音频文件
audio = AudioProcessor.load_audio(file_path)
# 转换为16kHz单声道PCM格式
pcm_audio = audio.set_frame_rate(16000).set_channels(1)
pcm_data = pcm_audio.raw_data
# 分块发送音频
chunk_size = 320 # 约100ms的音频数据
sequence_number = 0
async with websockets.connect(websocket_url) as websocket:
print(f"连接到WebSocket: {websocket_url}")
# 发送开始包
start_packet = create_websocket_packet(device_id, session_id, sequence_number, 1, b"")
await websocket.send(start_packet)
sequence_number += 1
print(f"发送开始包: 设备ID={device_id}, 会话ID={session_id}")
# 接收开始响应
try:
response = await asyncio.wait_for(websocket.recv(), timeout=5.0)
print(f"开始响应: {response}")
except asyncio.TimeoutError:
print("等待开始响应超时")
# 发送音频数据
for i in range(0, len(pcm_data), chunk_size):
if i>=3:
break
chunk = pcm_data[i:i+chunk_size]
audio_packet = create_websocket_packet(device_id, session_id, sequence_number, 0, chunk)
await websocket.send(audio_packet)
# 接收中间结果
try:
response = await asyncio.wait_for(websocket.recv(), timeout=1.0)
print(f"中间结果: {response}")
except asyncio.TimeoutError:
pass # 没有中间结果时继续
sequence_number += 1
# 添加延迟模拟实时流
await asyncio.sleep(0.1)
# 发送结束包
end_packet = create_websocket_packet(device_id, session_id, sequence_number, 2, b"")
await websocket.send(end_packet)
print("发送结束包")
# 接收最终结果
try:
response = await asyncio.wait_for(websocket.recv(), timeout=10.0)
print(f"最终结果: {response}")
except asyncio.TimeoutError:
print("等待最终结果超时")
print(f"音频文件 {file_path} 发送完成")
async def stream_large_audio():
asr = AliyunASR()
asr.device_id = "stream-device"
asr.session_id = "large-file-session"
# 创建音频流生成器
generator = audio_chunk_generator("./test/output(1).mp3", chunk_duration=500)
# generator1 = audio_chunk_generator("./test/16k.amr", chunk_duration=500)
# 转录MP3文件
print("开始转录MP3文件...")
async for result in transcribe_stream(asr, generator):
print(f"MP3中间结果: {result}")
async def send_audio_with_continuous_receive(file_path: str, websocket_url: str = "ws://118.25.13.180:8080/ws"):
"""
发送音频并持续接收所有返回结果
"""
device_id = "TalkingQ_10003B58862C"
session_id = "1000000189ABCDEFGHIJKLMNOPQRSTUV"
# 加载音频文件
# audio = AudioProcessor.load_audio(file_path)
with open(file_path, "rb") as f:
pcm_data = f.read()
# pcm_audio = audio.set_frame_rate(16000).set_channels(1)
# pcm_data = pcm_audio.raw_data
# pcm_data = AudioProcessor.audio_to_pcm(pcm_audio)
chunk_size = 320
sequence_number = 0
print(len(pcm_data))
async with websockets.connect(websocket_url) as websocket:
print(f"连接到WebSocket: {websocket_url}")
# 启动接收任务
receive_task = asyncio.create_task(receive_messages(websocket))
try:
# 发送开始包
start_packet = create_websocket_packet(device_id, session_id, sequence_number, 1, b"")
await websocket.send(start_packet)
sequence_number += 1
print("发送开始包")
# 发送音频数据
for i in range(0, len(pcm_data), chunk_size):
# if i>=1:
# break
chunk = pcm_data[i:i+chunk_size]
audio_packet = create_websocket_packet(device_id, session_id, sequence_number, 0, chunk)
await websocket.send(audio_packet)
sequence_number += 1
await asyncio.sleep(1)
# 发送结束包
end_packet = create_websocket_packet(device_id, session_id, sequence_number, 2, b"")
await websocket.send(end_packet)
print("发送结束包")
# 等待接收任务完成
await receive_task
except Exception as e:
print(f"发送过程中出错: {e}")
receive_task.cancel()
async def receive_messages(websocket):
"""
持续接收WebSocket消息
"""
try:
while True:
try:
message = await asyncio.wait_for(websocket.recv(), timeout=60.0)
print(f"收到消息: {message}")
except asyncio.TimeoutError:
print("接收超时,停止监听")
break
except websockets.exceptions.ConnectionClosed:
print("WebSocket连接已关闭")
break
except Exception as e:
print(f"接收消息时出错: {e}")
if __name__ == "__main__":
# 运行WebSocket测试 - 简单版本
# asyncio.run(send_audio_to_websocket("./test/output(1).mp3"))
# 运行WebSocket测试 - 持续接收版本
asyncio.run(send_audio_with_continuous_receive("./test/output(1).mp3"))
# 或者运行原来的测试
# asyncio.run(stream_large_audio())