add banbanmini backend
This commit is contained in:
0
talkingq-url/handlers/__init__.py
Normal file
0
talkingq-url/handlers/__init__.py
Normal file
87
talkingq-url/handlers/audio_packet_parser.py
Normal file
87
talkingq-url/handlers/audio_packet_parser.py
Normal 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
|
||||
181
talkingq-url/handlers/audio_session_handler.py
Normal file
181
talkingq-url/handlers/audio_session_handler.py
Normal 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, "已完成中断处理")
|
||||
80
talkingq-url/handlers/command_handler.py
Normal file
80
talkingq-url/handlers/command_handler.py
Normal 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
|
||||
135
talkingq-url/handlers/prompt_sound_handler.py
Normal file
135
talkingq-url/handlers/prompt_sound_handler.py
Normal 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)
|
||||
184
talkingq-url/handlers/response_coordinator.py
Normal file
184
talkingq-url/handlers/response_coordinator.py
Normal 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': '0', '1': '1', '2': '2', '3': '3', '4': '4',
|
||||
'5': '5', '6': '6', '7': '7', '8': '8', '9': '9',
|
||||
'A': 'a', 'B': 'b', 'C': 'c', 'D': 'd', 'E': 'e',
|
||||
'F': 'f', 'G': 'g', 'H': 'h', 'I': 'i', 'J': 'j',
|
||||
'K': 'k', 'L': 'l', 'M': 'm', 'N': 'n', 'O': 'o',
|
||||
'P': 'p', 'Q': 'q', 'R': 'r', 'S': 's', 'T': 't',
|
||||
'U': 'u', 'V': 'v', 'W': 'w', 'X': 'x', 'Y': 'y',
|
||||
'Z': 'z', 'a': 'a', 'b': 'b', 'c': 'c', 'd': 'd',
|
||||
'e': 'e', 'f': 'f', 'g': 'g', 'h': 'h', 'i': 'i',
|
||||
'j': 'j', 'k': 'k', 'l': 'l', 'm': 'm', 'n': 'n',
|
||||
'o': 'o', 'p': 'p', 'q': 'q', 'r': 'r', 's': 's',
|
||||
't': 't', 'u': 'u', 'v': 'v', 'w': 'w', 'x': 'x',
|
||||
'y': 'y', 'z': '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
|
||||
64
talkingq-url/handlers/response_processor.py
Normal file
64
talkingq-url/handlers/response_processor.py
Normal 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)
|
||||
185
talkingq-url/handlers/service_connection_handler.py
Normal file
185
talkingq-url/handlers/service_connection_handler.py
Normal 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 连接未建立",
|
||||
)
|
||||
53
talkingq-url/handlers/session_cleanup_handler.py
Normal file
53
talkingq-url/handlers/session_cleanup_handler.py
Normal 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], "已从会话管理器中移除会话"
|
||||
)
|
||||
75
talkingq-url/handlers/transcription_handler.py
Normal file
75
talkingq-url/handlers/transcription_handler.py
Normal 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} 秒"
|
||||
)
|
||||
66
talkingq-url/handlers/websocket_auth_handler.py
Normal file
66
talkingq-url/handlers/websocket_auth_handler.py
Normal 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
|
||||
55
talkingq-url/handlers/websocket_handler.py
Normal file
55
talkingq-url/handlers/websocket_handler.py
Normal 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连接开放,由客户端负责断开连接",
|
||||
)
|
||||
180
talkingq-url/handlers/websocket_message_handler.py
Normal file
180
talkingq-url/handlers/websocket_message_handler.py
Normal 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
|
||||
Reference in New Issue
Block a user