update:声纹识别对接的优化

This commit is contained in:
CGD
2025-07-08 11:25:54 +08:00
parent 3d64152dba
commit 5c83d63fe2
3 changed files with 175 additions and 144 deletions
+2 -129
View File
@@ -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))
+39 -15
View File
@@ -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