mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 17:43:55 +08:00
update:声纹识别对接的优化
This commit is contained in:
@@ -8,147 +8,20 @@ import threading
|
|||||||
import opuslib_next
|
import opuslib_next
|
||||||
import json
|
import json
|
||||||
import io
|
import io
|
||||||
import aiohttp
|
|
||||||
import time
|
import time
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from urllib.parse import urlparse, parse_qs
|
|
||||||
from typing import Optional, Tuple, List, Dict, Any
|
from typing import Optional, Tuple, List, Dict, Any
|
||||||
from core.handle.receiveAudioHandle import startToChat
|
from core.handle.receiveAudioHandle import startToChat
|
||||||
from core.handle.reportHandle import enqueue_asr_report
|
from core.handle.reportHandle import enqueue_asr_report
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
from core.handle.receiveAudioHandle import handleAudioMessage
|
from core.handle.receiveAudioHandle import handleAudioMessage
|
||||||
|
from core.utils.voiceprint_provider import VoiceprintProvider
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
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):
|
class ASRProviderBase(ABC):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
@@ -365,7 +238,7 @@ class ASRProviderBase(ABC):
|
|||||||
|
|
||||||
with wave.open(file_path, "wb") as wf:
|
with wave.open(file_path, "wb") as wf:
|
||||||
wf.setnchannels(1)
|
wf.setnchannels(1)
|
||||||
wf.setsampwidth(2)
|
wf.setsampwidth(2) # 2 bytes = 16-bit
|
||||||
wf.setframerate(16000)
|
wf.setframerate(16000)
|
||||||
wf.writeframes(b"".join(pcm_data))
|
wf.writeframes(b"".join(pcm_data))
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import uuid
|
import uuid
|
||||||
from typing import List, Dict
|
from typing import List, Dict
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
from config.settings import load_config
|
||||||
|
|
||||||
|
|
||||||
class Message:
|
class Message:
|
||||||
@@ -45,10 +46,9 @@ class Dialogue:
|
|||||||
dialogue.append({"role": m.role, "content": m.content})
|
dialogue.append({"role": m.role, "content": m.content})
|
||||||
|
|
||||||
def get_llm_dialogue(self) -> List[Dict[str, str]]:
|
def get_llm_dialogue(self) -> List[Dict[str, str]]:
|
||||||
dialogue = []
|
# 直接调用get_llm_dialogue_with_memory,传入None作为memory_str
|
||||||
for m in self.dialogue:
|
# 这样确保说话人功能在所有调用路径下都生效
|
||||||
self.getMessages(m, dialogue)
|
return self.get_llm_dialogue_with_memory(None)
|
||||||
return dialogue
|
|
||||||
|
|
||||||
def update_system_message(self, new_content: str):
|
def update_system_message(self, new_content: str):
|
||||||
"""更新或添加系统消息"""
|
"""更新或添加系统消息"""
|
||||||
@@ -62,10 +62,7 @@ class Dialogue:
|
|||||||
def get_llm_dialogue_with_memory(
|
def get_llm_dialogue_with_memory(
|
||||||
self, memory_str: str = None
|
self, memory_str: str = None
|
||||||
) -> List[Dict[str, str]]:
|
) -> List[Dict[str, str]]:
|
||||||
if memory_str is None or len(memory_str) == 0:
|
# 构建对话
|
||||||
return self.get_llm_dialogue()
|
|
||||||
|
|
||||||
# 构建带记忆的对话
|
|
||||||
dialogue = []
|
dialogue = []
|
||||||
|
|
||||||
# 添加系统提示和记忆
|
# 添加系统提示和记忆
|
||||||
@@ -74,16 +71,43 @@ class Dialogue:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if system_message:
|
if system_message:
|
||||||
# 构建增强的系统提示,包含说话人处理指导
|
# 基础系统提示
|
||||||
|
enhanced_system_prompt = system_message.content
|
||||||
|
|
||||||
|
# 添加说话人识别功能说明
|
||||||
speaker_guidance = "\n\n[说话人识别功能说明]\n" \
|
speaker_guidance = "\n\n[说话人识别功能说明]\n" \
|
||||||
"当用户消息包含 [说话人: 姓名] 前缀时,表示系统已识别出说话人身份。\n" \
|
"当用户消息为JSON格式包含speaker字段时(如:{\"speaker\": \"张三\", \"content\": \"消息内容\"}),表示系统已识别出说话人身份。\n" \
|
||||||
"请根据说话人的身份特征(如果之前有相关信息)来调整回应风格和内容。\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})
|
dialogue.append({"role": "system", "content": enhanced_system_prompt})
|
||||||
|
|
||||||
# 添加用户和助手的对话
|
# 添加用户和助手的对话
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user