diff --git a/main/xiaozhi-server/core/providers/asr/aliyun.py b/main/xiaozhi-server/core/providers/asr/aliyun.py index 74fb4091..5fa18253 100644 --- a/main/xiaozhi-server/core/providers/asr/aliyun.py +++ b/main/xiaozhi-server/core/providers/asr/aliyun.py @@ -3,8 +3,6 @@ import json import asyncio from typing import Optional, Tuple, List import opuslib_next -import wave -import io import os import uuid import hmac @@ -165,21 +163,6 @@ class ASRProvider(ASRProviderBase): request += "&enable_voice_detection=false" return request - def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str: - """PCM数据保存为WAV文件""" - module_name = __name__.split(".")[-1] - file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav" - file_path = os.path.join(self.output_dir, file_name) - - with wave.open(file_path, "wb") as wf: - wf.setnchannels(1) # 单声道 - wf.setsampwidth(2) # 16-bit - wf.setframerate(self.sample_rate) - wf.writeframes(b"".join(pcm_data)) - - logger.bind(tag=TAG).debug(f"音频文件已保存至: {file_path}") - return file_path - async def _send_request(self, pcm_data: bytes) -> Optional[str]: """发送请求到阿里云ASR服务""" try: diff --git a/main/xiaozhi-server/core/providers/asr/baidu.py b/main/xiaozhi-server/core/providers/asr/baidu.py index fe73fd45..87cb6a71 100644 --- a/main/xiaozhi-server/core/providers/asr/baidu.py +++ b/main/xiaozhi-server/core/providers/asr/baidu.py @@ -36,20 +36,6 @@ class ASRProvider(ASRProviderBase): # 确保输出目录存在 os.makedirs(self.output_dir, exist_ok=True) - def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str: - """PCM数据保存为WAV文件""" - module_name = __name__.split(".")[-1] - file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav" - file_path = os.path.join(self.output_dir, file_name) - - with wave.open(file_path, "wb") as wf: - wf.setnchannels(1) - wf.setsampwidth(2) # 2 bytes = 16-bit - wf.setframerate(16000) - wf.writeframes(b"".join(pcm_data)) - - return file_path - async def speech_to_text( self, opus_data: List[bytes], session_id: str ) -> Tuple[Optional[str], Optional[str]]: diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py index 9d974180..153f40e2 100644 --- a/main/xiaozhi-server/core/providers/asr/base.py +++ b/main/xiaozhi-server/core/providers/asr/base.py @@ -1,7 +1,10 @@ -from abc import ABC, abstractmethod -from typing import Optional, Tuple, List +import os +import uuid +import wave import opuslib_next +from abc import ABC, abstractmethod from config.logger import setup_logging +from typing import Optional, Tuple, List TAG = __name__ logger = setup_logging() @@ -11,10 +14,19 @@ class ASRProviderBase(ABC): def __init__(self): self.audio_format = "opus" - @abstractmethod def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str: """PCM数据保存为WAV文件""" - pass + module_name = __name__.split(".")[-1] + file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav" + file_path = os.path.join(self.output_dir, file_name) + + with wave.open(file_path, "wb") as wf: + wf.setnchannels(1) + wf.setsampwidth(2) # 2 bytes = 16-bit + wf.setframerate(16000) + wf.writeframes(b"".join(pcm_data)) + + return file_path @abstractmethod async def speech_to_text( diff --git a/main/xiaozhi-server/core/providers/asr/doubao.py b/main/xiaozhi-server/core/providers/asr/doubao.py index 403c181f..bd33af40 100644 --- a/main/xiaozhi-server/core/providers/asr/doubao.py +++ b/main/xiaozhi-server/core/providers/asr/doubao.py @@ -1,17 +1,13 @@ import time -import io -import wave import os -from typing import Optional, Tuple, List import uuid -import websockets import json import gzip - -import opuslib_next +import websockets +from config.logger import setup_logging +from typing import Optional, Tuple, List from core.providers.asr.base import ASRProviderBase -from config.logger import setup_logging TAG = __name__ logger = setup_logging() @@ -102,20 +98,6 @@ class ASRProvider(ASRProviderBase): # 确保输出目录存在 os.makedirs(self.output_dir, exist_ok=True) - def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str: - """PCM数据保存为WAV文件""" - module_name = __name__.split(".")[-1] - file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav" - file_path = os.path.join(self.output_dir, file_name) - - with wave.open(file_path, "wb") as wf: - wf.setnchannels(1) - wf.setsampwidth(2) # 2 bytes = 16-bit - wf.setframerate(16000) - wf.writeframes(b"".join(pcm_data)) - - return file_path - @staticmethod def _generate_header( message_type=CLIENT_FULL_REQUEST, message_type_specific_flags=NO_SEQUENCE diff --git a/main/xiaozhi-server/core/providers/asr/fun_local.py b/main/xiaozhi-server/core/providers/asr/fun_local.py index c8446574..e3b7fc83 100644 --- a/main/xiaozhi-server/core/providers/asr/fun_local.py +++ b/main/xiaozhi-server/core/providers/asr/fun_local.py @@ -53,20 +53,6 @@ class ASRProvider(ASRProviderBase): # device="cuda:0", # 启用GPU加速 ) - def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str: - """PCM数据保存为WAV文件""" - module_name = __name__.split(".")[-1] - file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav" - file_path = os.path.join(self.output_dir, file_name) - - with wave.open(file_path, "wb") as wf: - wf.setnchannels(1) - wf.setsampwidth(2) # 2 bytes = 16-bit - wf.setframerate(16000) - wf.writeframes(b"".join(pcm_data)) - - return file_path - async def speech_to_text( self, opus_data: List[bytes], session_id: str ) -> Tuple[Optional[str], Optional[str]]: diff --git a/main/xiaozhi-server/core/providers/asr/fun_server.py b/main/xiaozhi-server/core/providers/asr/fun_server.py index 97c9fdc9..ff411340 100644 --- a/main/xiaozhi-server/core/providers/asr/fun_server.py +++ b/main/xiaozhi-server/core/providers/asr/fun_server.py @@ -43,20 +43,6 @@ class ASRProvider(ASRProviderBase): self.ssl_context.check_hostname = False self.ssl_context.verify_mode = ssl.CERT_NONE - def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str: - """PCM数据保存为WAV文件""" - module_name = __name__.split(".")[-1] - file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav" - file_path = os.path.join(self.output_dir, file_name) - - with wave.open(file_path, "wb") as wf: - wf.setnchannels(1) - wf.setsampwidth(2) # 2 bytes = 16-bit - wf.setframerate(16000) - wf.writeframes(b"".join(pcm_data)) - - return file_path - async def _receive_responses(self, ws) -> None: """ Asynchronous generator to receive messages from the WebSocket. diff --git a/main/xiaozhi-server/core/providers/asr/sherpa_onnx_local.py b/main/xiaozhi-server/core/providers/asr/sherpa_onnx_local.py index 667e1606..24522d83 100644 --- a/main/xiaozhi-server/core/providers/asr/sherpa_onnx_local.py +++ b/main/xiaozhi-server/core/providers/asr/sherpa_onnx_local.py @@ -84,20 +84,6 @@ class ASRProvider(ASRProviderBase): use_itn=True, ) - def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str: - """PCM数据保存为WAV文件""" - module_name = __name__.split(".")[-1] - file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav" - file_path = os.path.join(self.output_dir, file_name) - - with wave.open(file_path, "wb") as wf: - wf.setnchannels(1) - wf.setsampwidth(2) # 2 bytes = 16-bit - wf.setframerate(16000) - wf.writeframes(b"".join(pcm_data)) - - return file_path - def read_wave(self, wave_filename: str) -> Tuple[np.ndarray, int]: """ Args: diff --git a/main/xiaozhi-server/core/providers/asr/tencent.py b/main/xiaozhi-server/core/providers/asr/tencent.py index 5ab1befb..22a7ed94 100644 --- a/main/xiaozhi-server/core/providers/asr/tencent.py +++ b/main/xiaozhi-server/core/providers/asr/tencent.py @@ -8,8 +8,6 @@ import os import uuid from typing import Optional, Tuple, List import wave -import opuslib_next - import requests from core.providers.asr.base import ASRProviderBase from config.logger import setup_logging @@ -33,20 +31,6 @@ class ASRProvider(ASRProviderBase): # 确保输出目录存在 os.makedirs(self.output_dir, exist_ok=True) - def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str: - """PCM数据保存为WAV文件""" - module_name = __name__.split(".")[-1] - file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav" - file_path = os.path.join(self.output_dir, file_name) - - with wave.open(file_path, "wb") as wf: - wf.setnchannels(1) - wf.setsampwidth(2) # 2 bytes = 16-bit - wf.setframerate(16000) - wf.writeframes(b"".join(pcm_data)) - - return file_path - async def speech_to_text( self, opus_data: List[bytes], session_id: str ) -> Tuple[Optional[str], Optional[str]]: