From 5c83d63fe2f2e0c8076611364899bff48c8a1f63 Mon Sep 17 00:00:00 2001 From: CGD <3030332422@qq.com> Date: Tue, 8 Jul 2025 11:25:54 +0800 Subject: [PATCH] =?UTF-8?q?update:=E5=A3=B0=E7=BA=B9=E8=AF=86=E5=88=AB?= =?UTF-8?q?=E5=AF=B9=E6=8E=A5=E7=9A=84=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../xiaozhi-server/core/providers/asr/base.py | 131 +---------------- main/xiaozhi-server/core/utils/dialogue.py | 54 +++++-- .../core/utils/voiceprint_provider.py | 134 ++++++++++++++++++ 3 files changed, 175 insertions(+), 144 deletions(-) create mode 100644 main/xiaozhi-server/core/utils/voiceprint_provider.py diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py index 8630f9b1..142e703f 100644 --- a/main/xiaozhi-server/core/providers/asr/base.py +++ b/main/xiaozhi-server/core/providers/asr/base.py @@ -8,147 +8,20 @@ import threading import opuslib_next import json import io -import aiohttp import time import concurrent.futures from abc import ABC, abstractmethod from config.logger import setup_logging -from urllib.parse import urlparse, parse_qs from typing import Optional, Tuple, List, Dict, Any from core.handle.receiveAudioHandle import startToChat from core.handle.reportHandle import enqueue_asr_report from core.utils.util import remove_punctuation_and_length from core.handle.receiveAudioHandle import handleAudioMessage +from core.utils.voiceprint_provider import VoiceprintProvider TAG = __name__ logger = setup_logging() -# 创建全局线程池执行器用于CPU密集型操作 -executor = concurrent.futures.ThreadPoolExecutor(max_workers=4) - -class VoiceprintProvider: - """声纹识别服务提供者""" - - def __init__(self, config: dict): - self.original_url = config.get("url", "") - self.speakers = config.get("speakers", []) - self.speaker_map = self._parse_speakers() - - # 解析API地址和密钥 - self.api_url = None - self.api_key = None - self.speaker_ids = [] - - if not self.original_url: - logger.bind(tag=TAG).warning("声纹识别URL未配置,声纹识别将被禁用") - self.enabled = False - else: - # 解析URL和key - parsed_url = urlparse(self.original_url) - base_url = f"{parsed_url.scheme}://{parsed_url.netloc}" - - # 从查询参数中提取key - query_params = parse_qs(parsed_url.query) - self.api_key = query_params.get('key', [''])[0] - - if not self.api_key: - logger.bind(tag=TAG).error("URL中未找到key参数,声纹识别将被禁用") - self.enabled = False - else: - # 构造identify接口地址 - self.api_url = f"{base_url}/voiceprint/identify" - - # 提取speaker_ids - for speaker_str in self.speakers: - try: - parts = speaker_str.split(",", 2) - if len(parts) >= 1: - speaker_id = parts[0].strip() - self.speaker_ids.append(speaker_id) - except Exception: - continue - - # 检查是否有有效的说话人配置 - if not self.speaker_ids: - logger.bind(tag=TAG).warning("未配置有效的说话人,声纹识别将被禁用") - self.enabled = False - else: - self.enabled = True - logger.bind(tag=TAG).info(f"声纹识别已配置: API={self.api_url}, 说话人={len(self.speaker_ids)}个") - - def _parse_speakers(self) -> Dict[str, Dict[str, str]]: - """解析说话人配置""" - speaker_map = {} - for speaker_str in self.speakers: - try: - parts = speaker_str.split(",", 2) - if len(parts) >= 3: - speaker_id, name, description = parts[0].strip(), parts[1].strip(), parts[2].strip() - speaker_map[speaker_id] = { - "name": name, - "description": description - } - except Exception as e: - logger.bind(tag=TAG).warning(f"解析说话人配置失败: {speaker_str}, 错误: {e}") - return speaker_map - - async def identify_speaker(self, audio_data: bytes, session_id: str) -> Optional[str]: - """识别说话人""" - if not self.enabled or not self.api_url or not self.api_key: - logger.bind(tag=TAG).debug("声纹识别功能已禁用或未配置,跳过识别") - return None - - try: - api_start_time = time.monotonic() - - # 准备请求头 - headers = { - 'Authorization': f'Bearer {self.api_key}', - 'Accept': 'application/json' - } - - # 准备multipart/form-data数据 - data = aiohttp.FormData() - data.add_field('speaker_ids', ','.join(self.speaker_ids)) - data.add_field('file', audio_data, filename='audio.wav', content_type='audio/wav') - - timeout = aiohttp.ClientTimeout(total=10) - - # 网络请求 - async with aiohttp.ClientSession(timeout=timeout) as session: - async with session.post(self.api_url, headers=headers, data=data) as response: - - if response.status == 200: - result = await response.json() - speaker_id = result.get("speaker_id") - score = result.get("score", 0) - total_elapsed_time = time.monotonic() - api_start_time - - logger.bind(tag=TAG).info(f"声纹识别耗时: {total_elapsed_time:.3f}s") - - # 置信度检查 - if score < 0.5: - logger.bind(tag=TAG).warning(f"声纹识别置信度较低: {score:.3f}") - - if speaker_id and speaker_id in self.speaker_map: - result_name = self.speaker_map[speaker_id]["name"] - return result_name - else: - logger.bind(tag=TAG).warning(f"未识别的说话人ID: {speaker_id}") - return "未知说话人" - else: - logger.bind(tag=TAG).error(f"声纹识别API错误: HTTP {response.status}") - return None - - except asyncio.TimeoutError: - elapsed = time.monotonic() - api_start_time - logger.bind(tag=TAG).error(f"声纹识别超时: {elapsed:.3f}s") - return None - except Exception as e: - elapsed = time.monotonic() - api_start_time - logger.bind(tag=TAG).error(f"声纹识别失败: {e}") - return None - class ASRProviderBase(ABC): def __init__(self): @@ -365,7 +238,7 @@ class ASRProviderBase(ABC): with wave.open(file_path, "wb") as wf: wf.setnchannels(1) - wf.setsampwidth(2) + wf.setsampwidth(2) # 2 bytes = 16-bit wf.setframerate(16000) wf.writeframes(b"".join(pcm_data)) diff --git a/main/xiaozhi-server/core/utils/dialogue.py b/main/xiaozhi-server/core/utils/dialogue.py index fbb1f7ad..69fb250e 100644 --- a/main/xiaozhi-server/core/utils/dialogue.py +++ b/main/xiaozhi-server/core/utils/dialogue.py @@ -1,6 +1,7 @@ import uuid from typing import List, Dict from datetime import datetime +from config.settings import load_config class Message: @@ -45,10 +46,9 @@ class Dialogue: dialogue.append({"role": m.role, "content": m.content}) def get_llm_dialogue(self) -> List[Dict[str, str]]: - dialogue = [] - for m in self.dialogue: - self.getMessages(m, dialogue) - return dialogue + # 直接调用get_llm_dialogue_with_memory,传入None作为memory_str + # 这样确保说话人功能在所有调用路径下都生效 + return self.get_llm_dialogue_with_memory(None) def update_system_message(self, new_content: str): """更新或添加系统消息""" @@ -62,10 +62,7 @@ class Dialogue: def get_llm_dialogue_with_memory( self, memory_str: str = None ) -> List[Dict[str, str]]: - if memory_str is None or len(memory_str) == 0: - return self.get_llm_dialogue() - - # 构建带记忆的对话 + # 构建对话 dialogue = [] # 添加系统提示和记忆 @@ -74,16 +71,43 @@ class Dialogue: ) if system_message: - # 构建增强的系统提示,包含说话人处理指导 + # 基础系统提示 + enhanced_system_prompt = system_message.content + + # 添加说话人识别功能说明 speaker_guidance = "\n\n[说话人识别功能说明]\n" \ - "当用户消息包含 [说话人: 姓名] 前缀时,表示系统已识别出说话人身份。\n" \ - "请根据说话人的身份特征(如果之前有相关信息)来调整回应风格和内容。\n" \ + "当用户消息为JSON格式包含speaker字段时(如:{\"speaker\": \"张三\", \"content\": \"消息内容\"}),表示系统已识别出说话人身份。\n" \ + "请根据说话人的身份特征来调整回应风格和内容。\n" \ "你可以称呼说话人的名字,并参考他们的特点进行个性化回应。" + enhanced_system_prompt += speaker_guidance + + # 添加说话人个性化描述 + try: + config = load_config() + voiceprint_config = config.get("plugins", {}).get("voiceprint", {}) + speakers = voiceprint_config.get("speakers", []) + + if speakers: + enhanced_system_prompt += "\n\n[已知说话人信息]" + for speaker_str in speakers: + try: + parts = speaker_str.split(",", 2) + if len(parts) >= 2: + speaker_id = parts[0].strip() + name = parts[1].strip() + # 如果描述为空,则为"" + description = parts[2].strip() if len(parts) >= 3 else "" + enhanced_system_prompt += f"\n- {name}:{description}" + except: + continue + except: + # 配置读取失败时忽略错误,不影响其他功能 + pass + + # 只有当有记忆时才添加记忆部分 + if memory_str and len(memory_str) > 0: + enhanced_system_prompt += f"\n\n以下是用户的历史记忆:\n```\n{memory_str}\n```" - enhanced_system_prompt = ( - f"{system_message.content}{speaker_guidance}\n\n" - f"以下是用户的历史记忆:\n```\n{memory_str}\n```" - ) dialogue.append({"role": "system", "content": enhanced_system_prompt}) # 添加用户和助手的对话 diff --git a/main/xiaozhi-server/core/utils/voiceprint_provider.py b/main/xiaozhi-server/core/utils/voiceprint_provider.py new file mode 100644 index 00000000..b241fb51 --- /dev/null +++ b/main/xiaozhi-server/core/utils/voiceprint_provider.py @@ -0,0 +1,134 @@ +import asyncio +import json +import time +import aiohttp +from urllib.parse import urlparse, parse_qs +from typing import Optional, Dict +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + + +class VoiceprintProvider: + """声纹识别服务提供者""" + + def __init__(self, config: dict): + self.original_url = config.get("url", "") + self.speakers = config.get("speakers", []) + self.speaker_map = self._parse_speakers() + + # 解析API地址和密钥 + self.api_url = None + self.api_key = None + self.speaker_ids = [] + + if not self.original_url: + logger.bind(tag=TAG).warning("声纹识别URL未配置,声纹识别将被禁用") + self.enabled = False + else: + # 解析URL和key + parsed_url = urlparse(self.original_url) + base_url = f"{parsed_url.scheme}://{parsed_url.netloc}" + + # 从查询参数中提取key + query_params = parse_qs(parsed_url.query) + self.api_key = query_params.get('key', [''])[0] + + if not self.api_key: + logger.bind(tag=TAG).error("URL中未找到key参数,声纹识别将被禁用") + self.enabled = False + else: + # 构造identify接口地址 + self.api_url = f"{base_url}/voiceprint/identify" + + # 提取speaker_ids + for speaker_str in self.speakers: + try: + parts = speaker_str.split(",", 2) + if len(parts) >= 1: + speaker_id = parts[0].strip() + self.speaker_ids.append(speaker_id) + except Exception: + continue + + # 检查是否有有效的说话人配置 + if not self.speaker_ids: + logger.bind(tag=TAG).warning("未配置有效的说话人,声纹识别将被禁用") + self.enabled = False + else: + self.enabled = True + logger.bind(tag=TAG).info(f"声纹识别已配置: API={self.api_url}, 说话人={len(self.speaker_ids)}个") + + def _parse_speakers(self) -> Dict[str, Dict[str, str]]: + """解析说话人配置""" + speaker_map = {} + for speaker_str in self.speakers: + try: + parts = speaker_str.split(",", 2) + if len(parts) >= 3: + speaker_id, name, description = parts[0].strip(), parts[1].strip(), parts[2].strip() + speaker_map[speaker_id] = { + "name": name, + "description": description + } + except Exception as e: + logger.bind(tag=TAG).warning(f"解析说话人配置失败: {speaker_str}, 错误: {e}") + return speaker_map + + async def identify_speaker(self, audio_data: bytes, session_id: str) -> Optional[str]: + """识别说话人""" + if not self.enabled or not self.api_url or not self.api_key: + logger.bind(tag=TAG).debug("声纹识别功能已禁用或未配置,跳过识别") + return None + + try: + api_start_time = time.monotonic() + + # 准备请求头 + headers = { + 'Authorization': f'Bearer {self.api_key}', + 'Accept': 'application/json' + } + + # 准备multipart/form-data数据 + data = aiohttp.FormData() + data.add_field('speaker_ids', ','.join(self.speaker_ids)) + data.add_field('file', audio_data, filename='audio.wav', content_type='audio/wav') + + timeout = aiohttp.ClientTimeout(total=10) + + # 网络请求 + async with aiohttp.ClientSession(timeout=timeout) as session: + async with session.post(self.api_url, headers=headers, data=data) as response: + + if response.status == 200: + result = await response.json() + speaker_id = result.get("speaker_id") + score = result.get("score", 0) + total_elapsed_time = time.monotonic() - api_start_time + + logger.bind(tag=TAG).info(f"声纹识别耗时: {total_elapsed_time:.3f}s") + + # 置信度检查 + if score < 0.5: + logger.bind(tag=TAG).warning(f"声纹识别置信度较低: {score:.3f}") + + if speaker_id and speaker_id in self.speaker_map: + result_name = self.speaker_map[speaker_id]["name"] + return result_name + else: + logger.bind(tag=TAG).warning(f"未识别的说话人ID: {speaker_id}") + return "未知说话人" + else: + logger.bind(tag=TAG).error(f"声纹识别API错误: HTTP {response.status}") + return None + + except asyncio.TimeoutError: + elapsed = time.monotonic() - api_start_time + logger.bind(tag=TAG).error(f"声纹识别超时: {elapsed:.3f}s") + return None + except Exception as e: + elapsed = time.monotonic() - api_start_time + logger.bind(tag=TAG).error(f"声纹识别失败: {e}") + return None