add banbanmini backend

This commit is contained in:
HycJack
2026-03-24 15:04:36 +08:00
parent 0d0f995dc2
commit 7510ca6df1
197 changed files with 13008 additions and 0 deletions

View File

View File

@@ -0,0 +1,87 @@
import struct
from fastapi import WebSocket
from services.audio_session import audio_session_manager, AudioSession
from utils.logger import session_logger
from services.factory import get_asr, get_llm, get_tts
DEVICE_ID_SIZE = 22 # 将大小设为22在实际解析中会更灵活处理
SESSION_ID_SIZE = 33
SAMPLE_RATE = 16000 # 固定为16K采样率
CHANNELS = 1 # 单声道
SAMPLE_WIDTH = 2 # 16位 = 2字节
async def parse_packet(data: bytes, websocket: WebSocket):
try:
min_header_size = 4 + 1 + 4 # sequence_number + packet_type + data_size
if len(data) < min_header_size:
session_logger.warning(
"unknown", "unknown", f"数据包太小: {len(data)} < {min_header_size}"
)
return None, None, None, None, None
null_pos = data.find(b'\x00')
if null_pos == -1 or null_pos > 25: # 如果没找到或超出最大长度限制
null_pos = 22 # 使用默认值
device_id = data[:null_pos].decode("ascii")
session_start = null_pos + 1
session_end = session_start + SESSION_ID_SIZE
session_data = data[session_start:session_end]
session_null_pos = session_data.find(b'\x00')
if session_null_pos != -1:
session_id = session_data[:session_null_pos].decode("ascii")
else:
session_id = session_data.decode("ascii").rstrip("\x00")
offset = session_start + len(session_id) + 1
if offset + min_header_size > len(data):
session_logger.warning(
"unknown", "unknown", f"解析剩余头信息时数据包太小: {len(data)} < {offset + min_header_size}"
)
return None, None, None, None, None
remaining_data = data[offset:]
offset = 0
sequence_number = struct.unpack("<I", remaining_data[offset:offset + 4])[0]
offset += 4
packet_type = remaining_data[offset]
offset += 1
data_size = struct.unpack("<I", remaining_data[offset:offset + 4])[0]
offset += 4
audio_data = remaining_data[offset:offset + data_size]
session_key = (device_id, session_id)
sample_rate = SAMPLE_RATE
session = await audio_session_manager.get_session(session_key)
if not session:
asr_service = await get_asr(None)
llm_service = await get_llm(None)
tts_service = await get_tts(None)
session = AudioSession(
asr_service=asr_service,
llm_service=llm_service,
tts_service=tts_service,
sample_rate=sample_rate,
)
session.device_id = device_id
session.session_id = session_id
await audio_session_manager.set_session(session_key, session)
if packet_type == 1:
session.sample_rate = sample_rate
session.last_received_seq = max(session.last_received_seq, sequence_number)
return session_key, session, packet_type, audio_data, sample_rate
except Exception as e:
session_logger.error(
"unknown",
"unknown",
f"解析WebSocket数据包时出错: {str(e)}",
exc_info=True,
)
return None, None, None, None, None

View File

@@ -0,0 +1,181 @@
import asyncio
import time
from fastapi import WebSocket
from services.audio_session import AudioSession
from services.task_manager import task_manager
from utils.logger import session_logger
from services.interrupt_handler import interrupt_handler
async def handle_websocket_data(
websocket: WebSocket,
session_key,
session,
packet_type,
audio_data,
sample_rate=16000,
):
if not session:
return
device_id, session_id = session_key
session.last_activity_time = asyncio.get_running_loop().time()
if packet_type == 1: # 开始包
try:
await websocket.send_text("START_ACK")
session_logger.info(device_id, session_id, "已发送开始包确认")
except Exception as e:
session_logger.error(device_id, session_id, f"发送开始包确认失败: {str(e)}")
session.sample_rate = sample_rate
await handle_start_packet(session, device_id, session_id)
session_logger.info(device_id, session_id, "语音识别开始使用16K采样率")
elif packet_type == 0: # 音频数据包
await handle_audio_packet(session, audio_data)
elif packet_type == 2: # 结束包
try:
await websocket.send_text("END_ACK")
session_logger.info(device_id, session_id, "已发送结束包确认")
except Exception as e:
session_logger.error(device_id, session_id, f"发送结束包确认失败: {str(e)}")
await handle_end_packet(session, session_key)
session_logger.info(device_id, session_id, "语音识别完成")
elif packet_type == 3: # 中断包
try:
await websocket.send_text("INTERRUPT_ACK")
session_logger.info(device_id, session_id, "已发送中断包确认")
except Exception as e:
session_logger.error(device_id, session_id, f"发送中断包确认失败: {str(e)}")
await handle_interrupt_packet(session, session_key)
session_logger.info(device_id, session_id, "接收到中断包")
async def handle_start_packet(session: AudioSession, device_id: str, session_id: str):
session_key = (device_id, session_id)
await interrupt_handler.set_interrupt_state(session_key, False)
session.start_time = time.perf_counter()
session_logger.info(device_id, session_id, "会话开始")
await interrupt_handler.register_session(session_key)
from handlers.command_handler import get_device_role
selected_role = await get_device_role(device_id)
session.send_queue = asyncio.Queue()
session.device_id = device_id
session.session_id = session_id
if hasattr(session, "send_task") and session.send_task:
session.send_task.cancel()
try:
await session.send_task
except asyncio.CancelledError:
pass
session.send_task = None
from handlers.service_connection_handler import prepare_asr_service
await prepare_asr_service(session, device_id, session_id)
from handlers.service_connection_handler import (
prepare_llm_service,
prepare_tts_service,
)
async def background_warmup():
try:
llm_task = await task_manager.create_task(
prepare_llm_service(session, device_id, session_id, selected_role),
device_id=device_id,
session_key=(device_id, session_id),
task_type="service_warmup"
)
tts_task = await task_manager.create_task(
prepare_tts_service(session, device_id, session_id, selected_role),
device_id=device_id,
session_key=(device_id, session_id),
task_type="service_warmup"
)
await asyncio.gather(llm_task, tts_task)
session_logger.info(device_id, session_id, "LLM和TTS服务预热完成")
except Exception as e:
session_logger.error(
device_id, session_id, f"服务预热过程中发生错误: {e}", exc_info=True
)
await task_manager.create_task(
background_warmup(),
device_id=device_id,
session_key=(device_id, session_id),
task_type="background_warmup"
)
session_logger.info(
device_id, session_id, "ASR服务已准备就绪后台继续预热其他服务"
)
async def handle_audio_packet(session: AudioSession, audio_data: bytes):
if session.asr_service.connected:
# if len(audio_data) == 0:
# print("Received empty audio data packet, ignoring.")
# return
# print("Sending audio data to ASR service...")
# print(f"Audio data length: {len(audio_data)} bytes")
# print(audio_data)
if session.send_queue.qsize() < 100: # 设置合理的上限
await session.send_queue.put(audio_data)
else:
session_logger.warning("unknown", "unknown", "音频队列过大,丢弃部分数据")
async def handle_end_packet(session: AudioSession, session_key: tuple):
device_id, session_id = session_key
session.end_time = time.perf_counter()
session_logger.info(device_id, session_id, "会话结束")
if session.send_task:
session.send_task.cancel()
try:
await session.send_task
except asyncio.CancelledError:
pass
from handlers.service_connection_handler import end_asr_session
from handlers.transcription_handler import handle_transcription
await end_asr_session(session)
await handle_transcription(session, device_id, session_id)
if session.llm_service and hasattr(session.llm_service, "close"):
await session.llm_service.close()
if session.tts_service:
await session.tts_service.close()
await interrupt_handler.remove_session(session_key)
async def handle_interrupt_packet(session: AudioSession, session_key: tuple):
"""处理中断包"""
device_id, session_id = session_key
session_logger.info(device_id, session_id, "处理中断包")
if not interrupt_handler.is_interrupted(session_key):
await interrupt_handler.set_interrupt_state(session_key, True)
if session.llm_service and hasattr(session.llm_service, "closed"):
session.llm_service.closed = True
is_processing = interrupt_handler.is_processing_interrupt(session_key)
if is_processing:
session_logger.info(
device_id, session_id, "检测到连续中断,中断请求将被排队"
)
await task_manager.create_task(
interrupt_handler.handle_interrupt(session_key),
device_id=device_id,
session_key=session_key,
task_type="interrupt"
)
if session.websocket and not is_processing:
await task_manager.create_task(
interrupt_handler.notify_client_interrupt_processed(session_key, session.websocket),
device_id=device_id,
session_key=session_key,
task_type="interrupt_notification"
)
else:
session_logger.info(
device_id, session_id, "会话已处于中断状态,重复的中断请求已被记录"
)
session_logger.info(device_id, session_id, "已完成中断处理")

View File

@@ -0,0 +1,80 @@
from config import settings
from services.device_config import DeviceConfig, device_config_manager
from services.role_manager import role_manager
from utils.logger import session_logger
async def update_device_config(device_id: str, new_role_key: str, language: str = None):
current_config = await device_config_manager.get_config(device_id, force_refresh=True)
preserved_language = language
if language is None and current_config and hasattr(current_config, 'preferred_language') and current_config.preferred_language:
preserved_language = current_config.preferred_language
session_logger.info(
device_id,
"config",
f"保留当前首选语言: {preserved_language}"
)
device_config = DeviceConfig(selected_role_key=new_role_key, preferred_language=preserved_language)
await device_config_manager.set_config(device_id, device_config)
session_logger.info(
device_id,
"config",
f"设备的角色已更新为 {new_role_key}" +
(f",首选语言: {preserved_language}" if preserved_language else "")
)
async def get_device_role(device_id: str, language: str = None) -> dict:
device_config = await device_config_manager.get_config(device_id, force_refresh=True)
preferred_language = None
if device_config and hasattr(device_config, 'preferred_language') and device_config.preferred_language:
preferred_language = device_config.preferred_language
session_logger.info(
device_id,
"config",
f"检测到设备首选语言配置: {preferred_language}"
)
effective_language = language or preferred_language
if device_config:
selected_role_key = device_config.selected_role_key
session_logger.info(
device_id,
"config",
f"使用配置的角色: {selected_role_key}" +
(f", 使用语言: {effective_language}" if effective_language else "")
)
else:
selected_role_key = settings.selected_role_key
session_logger.info(
device_id,
"config",
f"使用默认角色: {selected_role_key}" +
(f", 使用语言: {effective_language}" if effective_language else "")
)
selected_role = await role_manager.get_role_config_for_language(selected_role_key, effective_language)
if not selected_role:
all_roles = await role_manager.get_all_roles()
if all_roles:
first_role_key = next(iter(all_roles.keys()))
selected_role = await role_manager.get_role_config_for_language(first_role_key, effective_language)
session_logger.warning(
device_id,
"config",
f"未找到角色 {selected_role_key},使用默认角色: {first_role_key}"
)
else:
selected_role = {
"role_key": "default",
"name": "默认助手",
"content": "你是一个智能助手,请简洁回答问题。"
}
session_logger.warning(
device_id,
"config",
"未找到任何可用的角色配置,使用内置默认角色"
)
lang_info = effective_language or "默认"
session_logger.info(
device_id,
"config",
f"已选择角色配置: {selected_role.get('name', '未命名')}, 语言: {lang_info}"
)
return selected_role

View File

@@ -0,0 +1,135 @@
import os
from services.connection_manager import connection_manager
from services.device_config import device_config_manager
from services.role_manager import role_manager
from config import settings
from utils.logger import session_logger
async def handle_prompt_sound_request(
device_id: str, prompt_type: str, language: str = None
):
websocket = await connection_manager.get_connection(device_id)
if websocket and websocket.client_state.name == "CONNECTED":
try:
sound_files = {
"welcome": "welcome.mp3",
"error": "error.mp3",
"goodbye": "goodbye.mp3",
"interrupt": "interrupt.mp3",
"tts_error": "tts_error.mp3",
}
file_name = sound_files.get(prompt_type)
if not file_name:
session_logger.error(
device_id, "sound", f"未找到提示音类型: {prompt_type}"
)
return
device_config = await device_config_manager.get_config(device_id, force_refresh=True)
effective_language = language
if (
not effective_language
and device_config
and hasattr(device_config, "preferred_language")
and device_config.preferred_language
):
effective_language = device_config.preferred_language
session_logger.info(
device_id, "sound", f"使用设备首选语言选择提示音: {effective_language}"
)
selected_role_key = (
device_config.selected_role_key
if device_config
else settings.selected_role_key
)
session_logger.info(
device_id, "sound", f"从设备配置获取角色:{selected_role_key}"
)
selected_role = await role_manager.get_role_config_for_language(
selected_role_key, effective_language
)
if not selected_role:
session_logger.error(
device_id, "sound", f"未找到角色 {selected_role_key} 的配置"
)
return
base_url = None
if "multilingual" in selected_role:
current_lang = effective_language or selected_role.get(
"default_language", "zh"
)
if current_lang in selected_role["multilingual"]:
lang_config = selected_role["multilingual"][current_lang]
if "url" in lang_config:
base_url = lang_config["url"]
session_logger.info(
device_id, "sound", f"使用多语言({current_lang})配置的URL: {base_url}"
)
if not base_url and "url" in selected_role:
base_url = selected_role["url"]
session_logger.info(
device_id, "sound", f"使用角色基础URL: {base_url}"
)
if not base_url:
session_logger.error(
device_id, "sound", f"角色 {selected_role_key} 未配置URL"
)
return
sound_file_path = os.path.join(settings.assets_dir, base_url, file_name)
if not os.path.exists(sound_file_path):
session_logger.error(
device_id, "sound", f"提示音文件不存在: {sound_file_path}"
)
return
prompt_sound_url = f"http://{settings.server_host}:{settings.server_port}/assets/{base_url}/{file_name}"
try:
await websocket.send_text(f"PROMPT_SOUND_URL:{prompt_sound_url}")
session_logger.info(
device_id, "sound", f"已发送提示音URL给客户端: {prompt_sound_url}"
)
except Exception as e:
session_logger.error(
device_id, "sound", f"发送提示音URL失败: {e}"
)
except Exception as e:
session_logger.error(
device_id, "sound", f"准备提示音时出错: {e}", exc_info=True
)
else:
session_logger.error(
device_id, "sound", "未找到设备的 WebSocket 连接或连接已关闭"
)
async def send_welcome_sound(
device_id: str, role_key: str = None, language: str = None
):
device_config = await device_config_manager.get_config(device_id, force_refresh=True)
if role_key is None and device_config:
role_key = device_config.selected_role_key
session_logger.info(
device_id, "sound", f"使用设备配置的角色发送欢迎音效: {role_key}"
)
if (
language is None
and device_config
and hasattr(device_config, "preferred_language")
and device_config.preferred_language
):
language = device_config.preferred_language
session_logger.info(
device_id, "sound", f"使用设备首选语言发送欢迎音效: {language}"
)
await handle_prompt_sound_request(device_id, "welcome", language)

View File

@@ -0,0 +1,184 @@
import asyncio
from typing import Dict, List
from config import settings
from utils.logger import session_logger
from services.audio_session import AudioSession
from services.conversation_history import conversation_history_manager
from services.interrupt_handler import interrupt_handler
from services.task_manager import task_manager
from services.text_generator import TextGenerator
from services.tts_synthesizer import TTSSynthesizer
from services.audio_sender import AudioSender
from services.interruption_helper import InterruptionHelper
import re
async def _cleanup_queues(text_queue: asyncio.Queue, url_queue: asyncio.Queue, device_id: str, session_id: str):
"""清理响应处理中使用的队列"""
try:
# 清理text_queue
text_count = 0
while not text_queue.empty():
try:
item = text_queue.get_nowait()
text_queue.task_done()
text_count += 1
del item
except asyncio.QueueEmpty:
break
# 清理url_queue
url_count = 0
while not url_queue.empty():
try:
item = url_queue.get_nowait()
url_queue.task_done()
url_count += 1
del item
except asyncio.QueueEmpty:
break
session_logger.info(
device_id, session_id,
f"队列清理完成 - text_queue: {text_count}项, url_queue: {url_count}"
)
except Exception as e:
session_logger.error(
device_id, session_id, f"队列清理时出错: {e}"
)
SENSITIVE_RESPONSES = [
"你好,我无法给到相关内容",
"我无法给到相关内容",
"Sorry, I didn't catch that",
"I didn't catch that",
"让我们换一个话题",
"Let's change the topic",
"Changeons de sujet",
"Lassen Sie uns das Thema wechseln",
"Cambiemos de tema",
"Mari kita tukar topik"
]
def is_sensitive_response(reply: str) -> bool:
"""
检查回复是否为敏感话题回复,考虑不同语言的标点符号差异
"""
def normalize_text(text):
text = text.lower()
punctuations = r'[!"#$%&\'()*+,-./:;<=>?@\[\\\]^_`{|}~""''、。,!?:;()【】「」『』〈〉《》〔〕…—¥€£¥]'
text = re.sub(punctuations, '', text)
text = re.sub(r'\s+', '', text)
full_to_half = str.maketrans({
'': '0', '': '1', '': '2', '': '3', '': '4',
'': '5', '': '6', '': '7', '': '8', '': '9',
'': 'a', '': 'b', '': 'c', '': 'd', '': 'e',
'': 'f', '': 'g', '': 'h', '': 'i', '': 'j',
'': 'k', '': 'l', '': 'm', '': 'n', '': 'o',
'': 'p', '': 'q', '': 'r', '': 's', '': 't',
'': 'u', '': 'v', '': 'w', '': 'x', '': 'y',
'': 'z', '': 'a', '': 'b', '': 'c', '': 'd',
'': 'e', '': 'f', '': 'g', '': 'h', '': 'i',
'': 'j', '': 'k', '': 'l', '': 'm', '': 'n',
'': 'o', '': 'p', '': 'q', '': 'r', '': 's',
'': 't', '': 'u', '': 'v', '': 'w', '': 'x',
'': 'y', '': 'z', ' ': ''
})
text = text.translate(full_to_half)
return text
normalized_reply = normalize_text(reply)
for sensitive_reply in SENSITIVE_RESPONSES:
normalized_sensitive = normalize_text(sensitive_reply)
if normalized_sensitive in normalized_reply:
return True
return False
async def process_response(
transcript: str,
history: List[Dict[str, str]],
selected_role: dict,
device_id: str,
session_id: str,
device_history,
session: AudioSession,
language: str = None,
):
text_queue = asyncio.Queue() # 从LLM到TTS传递文本
url_queue = asyncio.Queue() # 从TTS到发送任务传递音频URL
session_key = (device_id, session_id)
await InterruptionHelper.register_queue_cleanup(session_key, url_queue)
await InterruptionHelper.register_queue_cleanup(session_key, text_queue)
text_generator = TextGenerator(device_id, session_id)
tts_synthesizer = TTSSynthesizer(device_id, session_id)
audio_sender = AudioSender(device_id, session_id, session.websocket)
llm_service = session.llm_service
tts_service = session.tts_service
if isinstance(tts_service.__class__.__name__, str) and tts_service.__class__.__name__ == "AliyunTTS":
tts_service = await get_tts(selected_role, language)
session.tts_service = tts_service
session_logger.warning(
device_id, session_id, "检测到AliyunTTS已不再支持自动切换到MiniMaxTTS"
)
text_gen_task = await task_manager.create_task(
text_generator.generate_text(
llm_service, transcript, history, selected_role, text_queue
),
device_id=device_id,
session_key=(device_id, session_id),
task_type="text_generation"
)
tts_task = await task_manager.create_task(
tts_synthesizer.synthesize_audio(
tts_service, text_queue, url_queue, selected_role, language # 传递语言参数到TTS
),
device_id=device_id,
session_key=(device_id, session_id),
task_type="tts_synthesis"
)
sender_task = await task_manager.create_task(
audio_sender.send_audio_urls(url_queue),
device_id=device_id,
session_key=(device_id, session_id),
task_type="audio_sender"
)
reply = ""
try:
reply = await text_gen_task
await asyncio.gather(tts_task, sender_task)
except Exception as e:
session_logger.error(
device_id, session_id, f"响应处理过程中出错: {e}", exc_info=True
)
finally:
# 清理队列
await _cleanup_queues(text_queue, url_queue, device_id, session_id)
if not interrupt_handler.is_interrupted(session_key):
if is_sensitive_response(reply):
session_logger.info(
device_id, session_id, "检测到敏感话题回复,该轮对话将不被添加到历史记录"
)
else:
device_history.history.append({"user": transcript, "assistant": reply})
if len(device_history.history) > settings.max_conversation_history:
device_history.history = device_history.history[-settings.max_conversation_history:]
role_key = selected_role.get("role_key", settings.selected_role_key)
device_history.role_key = role_key
await conversation_history_manager.set_history(device_id, device_history, role_key)
session_logger.info(
device_id,
session_id,
f"已将本轮对话添加到角色 {role_key} 的历史记录并保存到数据库"
)
return reply

View File

@@ -0,0 +1,64 @@
import asyncio
from typing import List, Dict
from services.audio_session import AudioSession
from handlers.prompt_sound_handler import handle_prompt_sound_request
from handlers.response_coordinator import process_response
from interfaces.llm import LLM
from interfaces.tts import TTS
from utils.logger import session_logger
from services.tts_error_manager import tts_error_manager
from implementations.minimax_tts import MiniMaxTTS
async def generate_and_process_response(
transcript: str,
history: List[Dict[str, str]],
selected_role: dict,
device_id: str,
session_id: str,
device_history,
session: AudioSession,
language: str = None, # 默认为None
):
llm_service: LLM = session.llm_service
if hasattr(llm_service, "device_id"):
llm_service.device_id = device_id
if hasattr(llm_service, "session_id"):
llm_service.session_id = session_id
tts_service: TTS = session.tts_service
if not isinstance(tts_service, MiniMaxTTS):
session_logger.warning(
device_id, session_id, "检测到TTS服务不是MiniMaxTTS正在切换到MiniMaxTTS"
)
tts_service = MiniMaxTTS(selected_role=selected_role)
session.tts_service = tts_service
try:
if isinstance(language, (int, float)) or (isinstance(language, str) and not language.isalpha()):
session_logger.warning(
device_id, session_id, f"检测到无效的语言代码: {language},将使用默认语言"
)
language = "zh" # 使用默认语言
await process_response(
transcript,
history,
selected_role,
device_id,
session_id,
device_history,
session,
language, # 传递合法化后的语言参数
)
except Exception as e:
session_logger.error(
device_id, session_id, f"生成或处理回复时出错: {e}", exc_info=True
)
if session.websocket and session.websocket.client_state.name == "CONNECTED":
try:
await session.websocket.send_text("TTS_ERROR")
except Exception as ws_error:
session_logger.error(device_id, session_id, f"发送TTS错误通知失败: {ws_error}")
await handle_prompt_sound_request(device_id, "tts_error")
finally:
await tts_error_manager.end_tts_session(device_id, session_id)

View File

@@ -0,0 +1,185 @@
import asyncio
from services.audio_session import AudioSession
from services.task_manager import task_manager
from utils.logger import session_logger
from services.interrupt_handler import interrupt_handler
import time
async def prepare_asr_service(session, device_id, session_id):
if hasattr(session.asr_service, "device_id"):
session.asr_service.device_id = device_id
if hasattr(session.asr_service, "session_id"):
session.asr_service.session_id = session_id
await establish_asr_connection(session)
async def prepare_llm_service(session, device_id, session_id, selected_role):
if hasattr(session.llm_service, "device_id"):
session.llm_service.device_id = device_id
if hasattr(session.llm_service, "session_id"):
session.llm_service.session_id = session_id
if hasattr(session.llm_service, "connect"):
await session.llm_service.connect()
from services.conversation_history import conversation_history_manager
from config import settings
device_history = await conversation_history_manager.get_history(device_id)
if device_history:
session.llm_service.history = device_history.history[
-settings.max_conversation_history :
]
async def prepare_tts_service(session, device_id, session_id, selected_role):
session.tts_service.selected_role = selected_role
if hasattr(session.tts_service, "device_id"):
session.tts_service.device_id = device_id
if hasattr(session.tts_service, "session_id"):
session.tts_service.session_id = session_id
await establish_tts_connection(session)
async def establish_asr_connection(session: AudioSession):
try:
if session.device_id and session.session_id:
session_key = (session.device_id, session.session_id)
if interrupt_handler.is_interrupted(session_key):
session_logger.info(
session.device_id, session.session_id, "会话已中断跳过建立ASR连接"
)
return
session.send_task = await task_manager.create_task(
send_audio_task(session),
device_id=session.device_id,
session_key=(session.device_id, session.session_id) if session.device_id and session.session_id else None,
task_type="audio_send"
)
await session.asr_service.connect()
if not session.asr_service.connected:
session_logger.error("unknown", "unknown", "语音识别 WebSocket 连接未建立")
return
session_logger.info(
"unknown", "unknown", f"语音识别连接已建立,采样率: {session.sample_rate}Hz"
)
except Exception as e:
session_logger.error(
"unknown", "unknown", f"建立语音识别 WebSocket 连接失败: {e}", exc_info=True
)
async def establish_tts_connection(session: AudioSession):
try:
if session.device_id and session.session_id:
session_key = (session.device_id, session.session_id)
if interrupt_handler.is_interrupted(session_key):
session_logger.info(
session.device_id, session.session_id, "会话已中断跳过建立TTS连接"
)
return
await session.tts_service.connect()
if not session.tts_service.connected:
session_logger.error(
session.device_id or "unknown",
session.session_id or "unknown",
"语音合成连接未建立",
)
return
session_logger.info(
session.device_id or "unknown",
session.session_id or "unknown",
"语音合成连接已建立",
)
except Exception as e:
session_logger.error(
session.device_id or "unknown",
session.session_id or "unknown",
f"建立语音合成连接失败: {e}",
)
async def send_audio_task(session: AudioSession):
try:
while True:
if session.device_id and session.session_id:
session_key = (session.device_id, session.session_id)
if interrupt_handler.is_interrupted(session_key):
session_logger.info(
session.device_id, session.session_id, "发送音频任务被中断"
)
break
audio_data = await session.send_queue.get()
if audio_data is None: # 结束信号
break
await session.asr_service.send_audio(audio_data)
except asyncio.CancelledError:
session_logger.info(
session.device_id or "unknown",
session.session_id or "unknown",
"发送音频数据任务被取消",
)
except Exception as e:
session_logger.error(
session.device_id or "unknown",
session.session_id or "unknown",
f"发送音频数据时出错: {str(e)}",
exc_info=True,
)
async def end_asr_session(session: AudioSession):
if session.asr_service.connected:
try:
asr_end_start_time = time.perf_counter()
await session.asr_service.send_end()
await session.asr_service.receive_results()
session.transcript = session.asr_service.transcript
asr_end_elapsed = time.perf_counter() - asr_end_start_time
session_logger.info(
session.device_id or "unknown",
session.session_id or "unknown",
f"ASR结束延迟: 从发送结束包到获取最终结果耗时 {asr_end_elapsed:.2f}",
)
except Exception as e:
session_logger.error(
session.device_id or "unknown",
session.session_id or "unknown",
f"ASR结束处理出错: {e}",
exc_info=True,
)
finally:
if session.send_task:
session.send_task.cancel()
try:
await asyncio.wait_for(session.send_task, timeout=2.0)
except (asyncio.CancelledError, asyncio.TimeoutError):
pass
except Exception as e:
session_logger.error(
session.device_id or "unknown",
session.session_id or "unknown",
f"取消发送任务时出错: {e}",
)
finally:
session.send_task = None
await session.asr_service.close()
session_logger.info(
session.device_id or "unknown",
session.session_id or "unknown",
"语音识别 WebSocket 连接和 ClientSession 已关闭",
)
else:
session_logger.error(
session.device_id or "unknown",
session.session_id or "unknown",
"语音识别 WebSocket 连接未建立",
)

View File

@@ -0,0 +1,53 @@
from services.interrupt_handler import interrupt_handler
from services.audio_session import audio_session_manager
from utils.logger import session_logger
async def handle_old_session_cleanup(old_session, session_key):
"""
清理旧会话资源
Args:
old_session: 需要清理的会话对象
session_key: 会话键 (device_id, session_id)
"""
try:
device_id, session_id = session_key
session_logger.info(device_id, session_id, "开始主动清理旧会话资源")
if hasattr(old_session, "register_interrupt_handlers"):
await old_session.register_interrupt_handlers()
await interrupt_handler.set_interrupt_state(session_key, True)
await interrupt_handler.handle_interrupt(session_key)
except Exception as e:
session_logger.error("unknown", "unknown", f"清理旧会话资源时出错: {e}")
async def cleanup_device_sessions(device_id):
"""
清理设备所有会话
Args:
device_id: 设备ID
"""
sessions_to_clean = []
for session_key, session in await audio_session_manager.get_all_sessions():
if session_key[0] == device_id:
sessions_to_clean.append((session_key, session))
for session_key, session in sessions_to_clean:
if hasattr(session, "set_interrupted"):
session_logger.info(
device_id, session_key[1], "WebSocket连接断开中断会话"
)
await session.set_interrupted(True)
await interrupt_handler.handle_interrupt(session_key)
else:
await interrupt_handler.set_interrupt_state(session_key, True)
if session.llm_service and hasattr(session.llm_service, "closed"):
session.llm_service.closed = True
for session_key, _ in sessions_to_clean:
await audio_session_manager.remove_session(session_key)
await interrupt_handler.remove_session(session_key)
session_logger.info(
device_id, session_key[1], "已从会话管理器中移除会话"
)

View File

@@ -0,0 +1,75 @@
import asyncio
import time
from services.audio_session import AudioSession
from utils.logger import session_logger
from config import settings
from services.device_config import device_config_manager
from utils.language_detector import LanguageDetector
async def handle_transcription(session: AudioSession, device_id: str, session_id: str):
start_time = time.perf_counter()
transcript = session.transcript
if transcript:
session_logger.info(device_id, session_id, f"最终转录结果: {transcript}")
detected_language = await LanguageDetector.detect_language(transcript)
session_logger.info(device_id, session_id, f"检测到的语言: {detected_language}")
device_config = await device_config_manager.get_config(device_id)
preferred_language = None
if device_config and hasattr(device_config, 'preferred_language') and device_config.preferred_language:
preferred_language = device_config.preferred_language
effective_language = detected_language or preferred_language
if detected_language and (not preferred_language or detected_language != preferred_language):
session_logger.info(
device_id,
session_id,
f"使用检测到的语言: {detected_language}"
)
elif preferred_language and detected_language != preferred_language:
session_logger.info(
device_id,
session_id,
f"检测到语言 {detected_language},但使用首选语言: {preferred_language}"
)
from services.conversation_history import (
DeviceConversationHistory,
conversation_history_manager,
)
from handlers.command_handler import get_device_role
selected_role = await get_device_role(device_id, effective_language)
role_key = selected_role.get("role_key") or settings.selected_role_key
device_history = await conversation_history_manager.get_history(device_id, role_key)
if not device_history:
device_history = DeviceConversationHistory()
await conversation_history_manager.set_history(device_id, device_history, role_key)
device_history.last_interaction_time = asyncio.get_running_loop().time()
history = device_history.history[-settings.max_conversation_history :]
from handlers.response_processor import generate_and_process_response
resp_start_time = time.perf_counter()
await generate_and_process_response(
transcript,
history,
selected_role,
device_id,
session_id,
device_history,
session,
effective_language,
)
resp_end_time = time.perf_counter()
session_logger.info(
device_id,
session_id,
f"回复生成和处理 耗时: {resp_end_time - resp_start_time:.2f}",
)
else:
session_logger.error("unknown", "unknown", "未收到转录结果")
from handlers.prompt_sound_handler import handle_prompt_sound_request
await handle_prompt_sound_request(device_id, "tts_error")
from services.tts_error_manager import tts_error_manager
await tts_error_manager.end_tts_session(device_id, session_id)
end_time = time.perf_counter()
session_logger.info(
device_id, session_id, f"总处理 耗时: {end_time - start_time:.2f}"
)

View File

@@ -0,0 +1,66 @@
import json
import asyncio
import time
from fastapi import WebSocket
from services.device_auth_manager import device_auth_manager
from services.connection_manager import connection_manager
from utils.logger import session_logger
async def authenticate_websocket(websocket: WebSocket):
"""
处理设备认证流程
Args:
websocket: WebSocket连接
Returns:
tuple: (认证状态, 设备ID) - (是否认证成功, 设备ID)
"""
device_id = None
authenticated = False
auth_timeout = 10 # 10秒认证超时
auth_start_time = time.time()
while not authenticated and time.time() - auth_start_time < auth_timeout:
try:
message = await asyncio.wait_for(
websocket.receive(),
timeout=auth_timeout - (time.time() - auth_start_time)
)
if message["type"] == "websocket.receive" and "text" in message:
try:
auth_data = json.loads(message["text"])
if "device_id" in auth_data and "serial_number" in auth_data:
device_id = auth_data["device_id"]
serial_number = auth_data["serial_number"]
authenticated = await device_auth_manager.authenticate_device(
device_id, serial_number
)
if authenticated:
await connection_manager.add_connection(device_id, websocket)
session_logger.info(
device_id, "auth", f"设备 {device_id} 认证成功"
)
await websocket.send_text(
json.dumps({"status": "authenticated"})
)
else:
session_logger.warning(
device_id, "auth", f"设备 {device_id} 认证失败"
)
await websocket.send_text(
json.dumps(
{"status": "error", "message": "Authentication failed"}
)
)
except json.JSONDecodeError:
session_logger.warning("unknown", "auth", "收到的不是有效的JSON认证消息")
elif message["type"] == "websocket.receive" and "bytes" in message:
pass
except asyncio.TimeoutError:
session_logger.warning("unknown", "auth", "WebSocket认证超时")
break
return authenticated, device_id

View File

@@ -0,0 +1,55 @@
from fastapi import WebSocket, WebSocketDisconnect
from utils.logger import session_logger
from services.connection_manager import connection_manager
from services.task_manager import task_manager
from handlers.websocket_auth_handler import authenticate_websocket
from handlers.websocket_message_handler import handle_websocket_messages
from handlers.session_cleanup_handler import cleanup_device_sessions
async def websocket_endpoint(websocket: WebSocket):
"""
WebSocket连接的主入口点
Args:
websocket: WebSocket连接
"""
await websocket.accept()
session_logger.info(
"unknown", "connection", f"WebSocket 连接已建立: {websocket.client}"
)
device_id = None
try:
authenticated, device_id = await authenticate_websocket(websocket)
if not authenticated:
session_logger.warning(
device_id or "unknown", "auth", "未认证的设备尝试连接,断开连接"
)
await websocket.send_text('{"status": "error", "message": "Not authenticated"}')
return
await handle_websocket_messages(websocket, device_id)
except WebSocketDisconnect:
session_logger.info(
device_id or "unknown",
"connection",
f"WebSocket 连接已关闭: {websocket.client}",
)
except Exception as e:
session_logger.error(
device_id or "unknown", "error", f"处理 WebSocket 数据时出错: {str(e)}"
)
finally:
if device_id:
await connection_manager.remove_connection(device_id)
await cleanup_device_sessions(device_id)
# 清理设备相关的所有异步任务
await task_manager.cancel_device_tasks(device_id)
session_logger.info(
device_id or "unknown",
"connection",
"服务端保持WebSocket连接开放由客户端负责断开连接",
)

View File

@@ -0,0 +1,180 @@
import json
import time
import asyncio
from fastapi import WebSocket
from handlers.audio_packet_parser import parse_packet
from handlers.audio_session_handler import handle_websocket_data
from services.audio_session import audio_session_manager
from services.interrupt_handler import interrupt_handler
from services.task_manager import task_manager
from utils.logger import session_logger
from handlers.prompt_sound_handler import handle_prompt_sound_request
from handlers.session_cleanup_handler import handle_old_session_cleanup
async def handle_websocket_messages(websocket: WebSocket, device_id: str):
"""
处理WebSocket连接中的所有消息
Args:
websocket: WebSocket连接
device_id: 设备ID
"""
current_active_session = None
first_audio_received_time = None
while True:
message = await websocket.receive()
if not websocket.client_state.name == "CONNECTED":
session_logger.info(
device_id,
"connection",
"WebSocket连接已关闭退出消息循环",
)
break
if message["type"] == "websocket.receive":
data = message["text"] if "text" in message else message["bytes"]
if isinstance(data, str):
await handle_text_message(websocket, device_id, data)
else: # 处理二进制音频数据
current_active_session, first_audio_received_time = await handle_binary_message(
websocket,
device_id,
data,
current_active_session,
first_audio_received_time
)
elif message["type"] == "websocket.disconnect":
break
async def handle_text_message(websocket: WebSocket, device_id: str, text_data: str):
"""处理文本消息"""
if text_data.startswith("REQUEST_PROMPT_SOUND:"):
prompt_type = text_data.split(":", 1)[1]
if device_id:
await handle_prompt_sound_request(device_id, prompt_type)
elif text_data == "NETWORK_RESET_ACKNOWLEDGED":
session_logger.info(device_id, "network", f"设备 {device_id} 确认网络重置")
elif text_data.startswith("FIRMWARE_UPDATE_STATUS:"):
status_info = text_data.split(":", 1)[1]
parts = status_info.split(",")
status_dict = {}
for part in parts:
if "=" in part:
key, value = part.split("=", 1)
status_dict[key.strip()] = value.strip()
from services.device_update_manager import device_firmware_update_manager
update_status = status_dict.get("status", "updating")
progress = float(status_dict.get("progress", "0")) if "progress" in status_dict else 0.0
await device_firmware_update_manager.update_firmware_progress(device_id, progress)
if update_status in ["success", "failed", "completed"]:
version = status_dict.get("version", "unknown")
await device_firmware_update_manager.update_firmware_update(
device_id,
firmware_version=version,
update_status=update_status
)
session_logger.info(device_id, "update", f"固件更新状态已保存: {update_status}, 进度: {progress}")
elif text_data.startswith("FIRMWARE_VERSION:"):
version_info = text_data.split(":", 1)[1].strip()
session_logger.info(device_id, "update", f"收到设备固件版本: {version_info}")
from services.device_update_manager import device_firmware_update_manager
await device_firmware_update_manager.update_device_firmware_version(
device_id,
version_info
)
await device_firmware_update_manager.update_firmware_update(
device_id,
version_info,
"success" # 收到版本信息说明设备当前状态正常
)
async def handle_binary_message(websocket: WebSocket, device_id: str, binary_data, current_active_session, first_audio_received_time):
"""处理二进制音频消息"""
session_key, session, packet_type, audio_data, sample_rate = await parse_packet(binary_data, websocket)
if session:
session.websocket = websocket
session_id = session_key[1]
if packet_type == 1: # 开始包
current_active_session = await handle_start_packet(
device_id, session_id, session_key, session, current_active_session
)
elif packet_type == 3: # 中断包
current_active_session = await handle_interrupt_packet(
device_id, session_id, session_key, session, current_active_session
)
elif packet_type == 2: # 结束包
if current_active_session == session_key:
current_active_session = None
elif packet_type == 0 and first_audio_received_time is None:
first_audio_received_time = time.perf_counter()
session.start_time = first_audio_received_time
session_logger.info(
device_id, session_id, f"收到第一个音频数据包时间: {first_audio_received_time}"
)
await handle_websocket_data(websocket, session_key, session, packet_type, audio_data, sample_rate)
return current_active_session, first_audio_received_time
async def handle_start_packet(device_id, session_id, session_key, session, current_active_session):
"""处理会话开始包"""
if current_active_session and current_active_session[1] != session_id:
old_session = await audio_session_manager.get_session(current_active_session)
if old_session:
session_logger.info(
device_id,
current_active_session[1],
f"检测到新会话 {session_id},优先中断旧会话 {current_active_session[1]}",
)
await interrupt_handler.set_interrupt_state(current_active_session, True)
if hasattr(old_session.llm_service, "close"):
await old_session.llm_service.close()
if hasattr(old_session.tts_service, "close"):
await old_session.tts_service.close()
await task_manager.create_task(
handle_old_session_cleanup(old_session, current_active_session),
device_id=device_id,
session_key=current_active_session,
task_type="cleanup"
)
current_active_session = session_key
await interrupt_handler.set_interrupt_state(session_key, False)
if hasattr(session.llm_service, "closed"):
session.llm_service.closed = False
session_logger.info(device_id, session_id, f"设置当前活跃会话: {session_id}")
await interrupt_handler.register_session(session_key)
return current_active_session
async def handle_interrupt_packet(device_id, session_id, session_key, session, current_active_session):
"""处理中断包"""
if current_active_session == session_key:
current_active_session = None
session_logger.info(device_id, session_id, "收到中断包,立即设置中断状态")
await interrupt_handler.set_interrupt_state(session_key, True)
if hasattr(session.llm_service, "closed"):
session.llm_service.closed = True
if interrupt_handler.is_processing_interrupt(session_key):
session_logger.info(device_id, session_id, "检测到连续中断,中断请求将被排队")
await task_manager.create_task(
interrupt_handler.handle_interrupt(session_key),
device_id=device_id,
session_key=session_key,
task_type="interrupt"
)
return current_active_session