From 98bfe863fe248a892e97c0286b27320477a2e526 Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Thu, 8 May 2025 15:08:11 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8Dtest=5Fpage.html=E6=B2=A1?= =?UTF-8?q?=E6=9C=89hello=E6=B6=88=E6=81=AF=EF=BC=8C=E5=AF=BC=E8=87=B4?= =?UTF-8?q?=E7=BC=BA=E5=A4=B1audio=5Fformat=E7=9A=84bug=20(#1151)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main/xiaozhi-server/core/connection.py | 9 ++- .../core/providers/asr/aliyun.py | 1 + .../core/providers/asr/baidu.py | 1 + .../xiaozhi-server/core/providers/asr/base.py | 11 ++- .../core/providers/asr/doubao.py | 1 + .../core/providers/asr/fun_local.py | 1 + .../core/providers/asr/fun_server.py | 81 ++++++++++++------- .../core/providers/asr/sherpa_onnx_local.py | 1 + .../core/providers/asr/tencent.py | 74 ++++++++++------- 9 files changed, 116 insertions(+), 64 deletions(-) diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 4c355b5b..ff8c8b4f 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -241,9 +241,12 @@ class ConnectionHandler: def _initialize_components(self): """初始化组件""" - self.prompt = self.config["prompt"] - self.change_system_prompt(self.prompt) - self.logger.bind(tag=TAG).info(f"初始化组件: prompt成功 {self.prompt[:50]}...") + if self.config.get("prompt") is not None: + self.prompt = self.config["prompt"] + self.change_system_prompt(self.prompt) + self.logger.bind(tag=TAG).info( + f"初始化组件: prompt成功 {self.prompt[:50]}..." + ) """初始化本地组件""" if self.vad is None: diff --git a/main/xiaozhi-server/core/providers/asr/aliyun.py b/main/xiaozhi-server/core/providers/asr/aliyun.py index 14421d7d..6606168c 100644 --- a/main/xiaozhi-server/core/providers/asr/aliyun.py +++ b/main/xiaozhi-server/core/providers/asr/aliyun.py @@ -91,6 +91,7 @@ class AccessToken: class ASRProvider(ASRProviderBase): def __init__(self, config: dict, delete_audio_file: bool): + super().__init__() """阿里云ASR初始化""" # 新增空值判断逻辑 self.access_key_id = config.get("access_key_id") diff --git a/main/xiaozhi-server/core/providers/asr/baidu.py b/main/xiaozhi-server/core/providers/asr/baidu.py index 58b26375..42e14d8b 100644 --- a/main/xiaozhi-server/core/providers/asr/baidu.py +++ b/main/xiaozhi-server/core/providers/asr/baidu.py @@ -20,6 +20,7 @@ logger = setup_logging() class ASRProvider(ASRProviderBase): def __init__(self, config: dict, delete_audio_file: bool = True): + super().__init__() self.app_id = config.get("app_id") self.api_key = config.get("api_key") self.secret_key = config.get("secret_key") diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py index 61c1b754..227a906d 100644 --- a/main/xiaozhi-server/core/providers/asr/base.py +++ b/main/xiaozhi-server/core/providers/asr/base.py @@ -8,16 +8,21 @@ logger = setup_logging() 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 @abstractmethod - async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]: + async def speech_to_text( + self, opus_data: List[bytes], session_id: str + ) -> Tuple[Optional[str], Optional[str]]: """将语音数据转换为文本""" pass - + def set_audio_format(self, format: str) -> None: """设置音频格式""" self.audio_format = format @@ -36,4 +41,4 @@ class ASRProviderBase(ABC): except opuslib_next.OpusError as e: logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True) - return pcm_data \ No newline at end of file + return pcm_data diff --git a/main/xiaozhi-server/core/providers/asr/doubao.py b/main/xiaozhi-server/core/providers/asr/doubao.py index 18e64a9b..403c181f 100644 --- a/main/xiaozhi-server/core/providers/asr/doubao.py +++ b/main/xiaozhi-server/core/providers/asr/doubao.py @@ -85,6 +85,7 @@ def parse_response(res): class ASRProvider(ASRProviderBase): def __init__(self, config: dict, delete_audio_file: bool): + super().__init__() self.appid = config.get("appid") self.cluster = config.get("cluster") self.access_token = config.get("access_token") diff --git a/main/xiaozhi-server/core/providers/asr/fun_local.py b/main/xiaozhi-server/core/providers/asr/fun_local.py index 107417b9..ec3df17f 100644 --- a/main/xiaozhi-server/core/providers/asr/fun_local.py +++ b/main/xiaozhi-server/core/providers/asr/fun_local.py @@ -33,6 +33,7 @@ class CaptureOutput: class ASRProvider(ASRProviderBase): def __init__(self, config: dict, delete_audio_file: bool): + super().__init__() self.model_dir = config.get("model_dir") self.output_dir = config.get("output_dir") # 修正配置键名 self.delete_audio_file = delete_audio_file diff --git a/main/xiaozhi-server/core/providers/asr/fun_server.py b/main/xiaozhi-server/core/providers/asr/fun_server.py index f20d5be7..53eb4e7f 100644 --- a/main/xiaozhi-server/core/providers/asr/fun_server.py +++ b/main/xiaozhi-server/core/providers/asr/fun_server.py @@ -9,22 +9,29 @@ import wave import websockets from config.logger import setup_logging import asyncio + TAG = __name__ logger = setup_logging() + class ASRProvider(ASRProviderBase): def __init__(self, config: dict, delete_audio_file: bool): - ''' + """ Initialize the ASRProvider with server configuration. :param config: Dictionary containing 'host', 'port', and 'is_ssl'. :param delete_audio_file: Boolean to indicate whether to delete audio files after processing. - ''' - self.host = config.get('host', 'localhost') - self.port = config.get('port', 10095) - self.is_ssl = config.get('is_ssl', True) + """ + super().__init__() + self.host = config.get("host", "localhost") + self.port = config.get("port", 10095) + self.is_ssl = config.get("is_ssl", True) self.output_dir = config.get("output_dir") self.delete_audio_file = delete_audio_file - self.uri = f"wss://{self.host}:{self.port}" if self.is_ssl else f"ws://{self.host}:{self.port}" + self.uri = ( + f"wss://{self.host}:{self.port}" + if self.is_ssl + else f"ws://{self.host}:{self.port}" + ) self.ssl_context = ssl.SSLContext() if self.is_ssl else None if self.ssl_context: self.ssl_context.check_hostname = False @@ -45,10 +52,10 @@ class ASRProvider(ASRProviderBase): return file_path async def _receive_responses(self, ws) -> None: - ''' + """ Asynchronous generator to receive messages from the WebSocket. Yields each message as it is received. - ''' + """ text = "" while True: try: @@ -61,30 +68,35 @@ class ASRProvider(ASRProviderBase): else: text += response_data.get("text", "") except asyncio.TimeoutError: - logger.bind(tag=TAG).error("Timeout while waiting for response from WebSocket.") + logger.bind(tag=TAG).error( + "Timeout while waiting for response from WebSocket." + ) break except websockets.exceptions.ConnectionClosed as e: logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}") break return text + async def _send_data(self, ws, pcm_data: bytes, session_id: str) -> tuple: - ''' + """ Internal method to handle WebSocket communication. Reuses the persistent WebSocket connection if available. :param pcm_data: PCM audio data to send. :param session_id: Unique session identifier. :return: Tuple containing recognized text and optional timestamp. - ''' + """ # Send initial configuration message - config_message = json.dumps({ - "mode": "offline", - "chunk_size": [5, 10, 5], - "chunk_interval": 10, - "wav_name": session_id, - "is_speaking": True, - "itn": False - }) + config_message = json.dumps( + { + "mode": "offline", + "chunk_size": [5, 10, 5], + "chunk_interval": 10, + "wav_name": session_id, + "is_speaking": True, + "itn": False, + } + ) await ws.send(config_message) logger.bind(tag=TAG).debug(f"Sent configuration message: {config_message}") @@ -97,14 +109,15 @@ class ASRProvider(ASRProviderBase): await ws.send(end_message) logger.bind(tag=TAG).debug(f"Sent end message: {end_message}") - - async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]: - ''' + async def speech_to_text( + self, opus_data: List[bytes], session_id: str + ) -> Tuple[Optional[str], Optional[str]]: + """ Convert speech data to text using FunASR. :param opus_data: List of Opus-encoded audio data chunks. :param session_id: Unique session identifier. :return: Tuple containing recognized text and optional timestamp. - ''' + """ file_path = None if self.audio_format == "pcm": pcm_data = opus_data @@ -118,16 +131,19 @@ class ASRProvider(ASRProviderBase): else: file_path = self.save_audio_to_file(pcm_data, session_id) - async with websockets.connect(self.uri, subprotocols=["binary"], ping_interval=None, ssl=self.ssl_context) as ws: + async with websockets.connect( + self.uri, subprotocols=["binary"], ping_interval=None, ssl=self.ssl_context + ) as ws: try: # Use asyncio to handle WebSocket communication - send_task = asyncio.create_task(self._send_data(ws, combined_pcm_data, session_id)) + send_task = asyncio.create_task( + self._send_data(ws, combined_pcm_data, session_id) + ) receive_task = asyncio.create_task(self._receive_responses(ws)) # Gather tasks with error handling done, pending = await asyncio.wait( - [send_task, receive_task], - return_when=asyncio.FIRST_EXCEPTION + [send_task, receive_task], return_when=asyncio.FIRST_EXCEPTION ) # Cancel any pending tasks @@ -141,11 +157,16 @@ class ASRProvider(ASRProviderBase): # Get the result from the receive task result = receive_task.result() - return result, file_path # Return the recognized text and timestamp (if any) + return ( + result, + file_path, + ) # Return the recognized text and timestamp (if any) except websockets.exceptions.ConnectionClosed as e: logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}") return "", file_path except Exception as e: - logger.bind(tag=TAG).error(f"Error during speech-to-text conversion: {e}", exc_info=True) - return "", file_path \ No newline at end of file + logger.bind(tag=TAG).error( + f"Error during speech-to-text conversion: {e}", exc_info=True + ) + return "", file_path 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 2b86f6d7..667e1606 100644 --- a/main/xiaozhi-server/core/providers/asr/sherpa_onnx_local.py +++ b/main/xiaozhi-server/core/providers/asr/sherpa_onnx_local.py @@ -37,6 +37,7 @@ class CaptureOutput: class ASRProvider(ASRProviderBase): def __init__(self, config: dict, delete_audio_file: bool): + super().__init__() self.model_dir = config.get("model_dir") self.output_dir = config.get("output_dir") self.delete_audio_file = delete_audio_file diff --git a/main/xiaozhi-server/core/providers/asr/tencent.py b/main/xiaozhi-server/core/providers/asr/tencent.py index 2aa7bec5..d2db875e 100644 --- a/main/xiaozhi-server/core/providers/asr/tencent.py +++ b/main/xiaozhi-server/core/providers/asr/tencent.py @@ -17,12 +17,14 @@ from config.logger import setup_logging TAG = __name__ logger = setup_logging() + class ASRProvider(ASRProviderBase): API_URL = "https://asr.tencentcloudapi.com" API_VERSION = "2019-06-14" FORMAT = "pcm" # 支持的音频格式:pcm, wav, mp3 def __init__(self, config: dict, delete_audio_file: bool = True): + super().__init__() self.secret_id = config.get("secret_id") self.secret_key = config.get("secret_key") self.output_dir = config.get("output_dir") @@ -45,7 +47,9 @@ class ASRProvider(ASRProviderBase): return file_path - async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]: + async def speech_to_text( + self, opus_data: List[bytes], session_id: str + ) -> Tuple[Optional[str], Optional[str]]: """将语音数据转换为文本""" if not opus_data: logger.bind(tag=TAG).warn("音频数据为空!") @@ -72,7 +76,7 @@ class ASRProvider(ASRProviderBase): self.save_audio_to_file(pcm_data, session_id) # 将音频数据转换为Base64编码 - base64_audio = base64.b64encode(combined_pcm_data).decode('utf-8') + base64_audio = base64.b64encode(combined_pcm_data).decode("utf-8") # 构建请求体 request_body = self._build_request_body(base64_audio) @@ -85,7 +89,9 @@ class ASRProvider(ASRProviderBase): result = self._send_request(request_body, timestamp, authorization) if result: - logger.bind(tag=TAG).debug(f"腾讯云语音识别耗时: {time.time() - start_time:.3f}s | 结果: {result}") + logger.bind(tag=TAG).debug( + f"腾讯云语音识别耗时: {time.time() - start_time:.3f}s | 结果: {result}" + ) return result, file_path @@ -102,7 +108,7 @@ class ASRProvider(ASRProviderBase): "SourceType": 1, # 音频数据来源为语音文件 "VoiceFormat": self.FORMAT, # 音频格式 "Data": base64_audio, # Base64编码的音频数据 - "DataLen": len(base64_audio) # 数据长度 + "DataLen": len(base64_audio), # 数据长度 } return json.dumps(request_map) @@ -135,9 +141,11 @@ class ASRProvider(ASRProviderBase): action = "SentenceRecognition" # 接口名称 # 构建规范头部信息,注意顺序和格式 - canonical_headers = f"content-type:{content_type.lower()}\n" + \ - f"host:{host.lower()}\n" + \ - f"x-tc-action:{action.lower()}\n" + canonical_headers = ( + f"content-type:{content_type.lower()}\n" + + f"host:{host.lower()}\n" + + f"x-tc-action:{action.lower()}\n" + ) signed_headers = "content-type;host;x-tc-action" @@ -145,21 +153,25 @@ class ASRProvider(ASRProviderBase): payload_hash = self._sha256_hex(request_body) # 构建规范请求字符串 - canonical_request = f"{http_request_method}\n" + \ - f"{canonical_uri}\n" + \ - f"{canonical_query_string}\n" + \ - f"{canonical_headers}\n" + \ - f"{signed_headers}\n" + \ - f"{payload_hash}" + canonical_request = ( + f"{http_request_method}\n" + + f"{canonical_uri}\n" + + f"{canonical_query_string}\n" + + f"{canonical_headers}\n" + + f"{signed_headers}\n" + + f"{payload_hash}" + ) # 计算规范请求的哈希值 hashed_canonical_request = self._sha256_hex(canonical_request) # 构建待签名字符串 - string_to_sign = f"{algorithm}\n" + \ - f"{timestamp}\n" + \ - f"{credential_scope}\n" + \ - f"{hashed_canonical_request}" + string_to_sign = ( + f"{algorithm}\n" + + f"{timestamp}\n" + + f"{credential_scope}\n" + + f"{hashed_canonical_request}" + ) # 计算签名密钥 secret_date = self._hmac_sha256(f"TC3{self.secret_key}", date) @@ -167,13 +179,17 @@ class ASRProvider(ASRProviderBase): secret_signing = self._hmac_sha256(secret_service, "tc3_request") # 计算签名 - signature = self._bytes_to_hex(self._hmac_sha256(secret_signing, string_to_sign)) + signature = self._bytes_to_hex( + self._hmac_sha256(secret_signing, string_to_sign) + ) # 构建授权头 - authorization = f"{algorithm} " + \ - f"Credential={self.secret_id}/{credential_scope}, " + \ - f"SignedHeaders={signed_headers}, " + \ - f"Signature={signature}" + authorization = ( + f"{algorithm} " + + f"Credential={self.secret_id}/{credential_scope}, " + + f"SignedHeaders={signed_headers}, " + + f"Signature={signature}" + ) return timestamp, authorization @@ -181,7 +197,9 @@ class ASRProvider(ASRProviderBase): logger.bind(tag=TAG).error(f"生成认证头失败: {e}", exc_info=True) raise RuntimeError(f"生成认证头失败: {e}") - def _send_request(self, request_body: str, timestamp: str, authorization: str) -> Optional[str]: + def _send_request( + self, request_body: str, timestamp: str, authorization: str + ) -> Optional[str]: """发送请求到腾讯云API""" headers = { "Content-Type": "application/json; charset=utf-8", @@ -190,7 +208,7 @@ class ASRProvider(ASRProviderBase): "X-TC-Action": "SentenceRecognition", "X-TC-Version": self.API_VERSION, "X-TC-Timestamp": timestamp, - "X-TC-Region": "ap-shanghai" + "X-TC-Region": "ap-shanghai", } try: @@ -221,16 +239,16 @@ class ASRProvider(ASRProviderBase): def _sha256_hex(self, data: str) -> str: """计算字符串的SHA256哈希值""" - digest = hashlib.sha256(data.encode('utf-8')).digest() + digest = hashlib.sha256(data.encode("utf-8")).digest() return self._bytes_to_hex(digest) def _hmac_sha256(self, key, data: str) -> bytes: """计算HMAC-SHA256""" if isinstance(key, str): - key = key.encode('utf-8') + key = key.encode("utf-8") - return hmac.new(key, data.encode('utf-8'), hashlib.sha256).digest() + return hmac.new(key, data.encode("utf-8"), hashlib.sha256).digest() def _bytes_to_hex(self, bytes_data: bytes) -> str: """字节数组转十六进制字符串""" - return ''.join(f"{b:02x}" for b in bytes_data) \ No newline at end of file + return "".join(f"{b:02x}" for b in bytes_data)