update:修复退出保存记忆bug (#619)

This commit is contained in:
hrz
2025-04-01 00:03:05 +08:00
committed by GitHub
parent 50490b70b7
commit 9430094072
2 changed files with 68 additions and 36 deletions
+1 -5
View File
@@ -195,7 +195,7 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).error(f"Connection error: {str(e)}-{stack_trace}") self.logger.bind(tag=TAG).error(f"Connection error: {str(e)}-{stack_trace}")
return return
finally: finally:
self._save_and_close(ws) await self._save_and_close(ws)
async def _save_and_close(self, ws): async def _save_and_close(self, ws):
"""保存记忆并关闭连接""" """保存记忆并关闭连接"""
@@ -253,13 +253,9 @@ class ConnectionHandler:
else: else:
# 否则使用主LLM # 否则使用主LLM
self.intent.set_llm(self.llm) self.intent.set_llm(self.llm)
self.logger.bind(tag=TAG).info("意图识别使用主LLM")
# 记录意图识别LLM初始化耗时 # 记录意图识别LLM初始化耗时
intent_llm_init_time = time.time() - intent_llm_init_start intent_llm_init_time = time.time() - intent_llm_init_start
self.logger.bind(tag=TAG).info(
f"意图识别LLM初始化完成,耗时: {intent_llm_init_time:.4f}"
)
"""加载位置信息""" """加载位置信息"""
self.client_ip_info = get_ip_info(self.client_ip) self.client_ip_info = get_ip_info(self.client_ip)
+67 -31
View File
@@ -12,50 +12,72 @@ class WebSocketServer:
def __init__(self, config: dict): def __init__(self, config: dict):
self.config = config self.config = config
self.logger = setup_logging() self.logger = setup_logging()
self._vad, self._asr, self._llm, self._tts, self._memory, self.intent = self._create_processing_instances() self._vad, self._asr, self._llm, self._tts, self._memory, self.intent = (
self._create_processing_instances()
)
self.active_connections = set() # 添加全局连接记录 self.active_connections = set() # 添加全局连接记录
def _create_processing_instances(self): def _create_processing_instances(self):
memory_cls_name = self.config["selected_module"].get("Memory", "nomem") # 默认使用nomem memory_cls_name = self.config["selected_module"].get(
has_memory_cfg = self.config.get("Memory") and memory_cls_name in self.config["Memory"] "Memory", "nomem"
) # 默认使用nomem
has_memory_cfg = (
self.config.get("Memory") and memory_cls_name in self.config["Memory"]
)
memory_cfg = self.config["Memory"][memory_cls_name] if has_memory_cfg else {} memory_cfg = self.config["Memory"][memory_cls_name] if has_memory_cfg else {}
"""创建处理模块实例""" """创建处理模块实例"""
return ( return (
vad.create_instance( vad.create_instance(
self.config["selected_module"]["VAD"], self.config["selected_module"]["VAD"],
self.config["VAD"][self.config["selected_module"]["VAD"]] self.config["VAD"][self.config["selected_module"]["VAD"]],
), ),
asr.create_instance( asr.create_instance(
self.config["selected_module"]["ASR"] (
if not 'type' in self.config["ASR"][self.config["selected_module"]["ASR"]] self.config["selected_module"]["ASR"]
else if not "type"
self.config["ASR"][self.config["selected_module"]["ASR"]]["type"], in self.config["ASR"][self.config["selected_module"]["ASR"]]
else self.config["ASR"][self.config["selected_module"]["ASR"]][
"type"
]
),
self.config["ASR"][self.config["selected_module"]["ASR"]], self.config["ASR"][self.config["selected_module"]["ASR"]],
self.config["delete_audio"] self.config["delete_audio"],
), ),
llm.create_instance( llm.create_instance(
self.config["selected_module"]["LLM"] (
if not 'type' in self.config["LLM"][self.config["selected_module"]["LLM"]] self.config["selected_module"]["LLM"]
else if not "type"
self.config["LLM"][self.config["selected_module"]["LLM"]]['type'], in self.config["LLM"][self.config["selected_module"]["LLM"]]
else self.config["LLM"][self.config["selected_module"]["LLM"]][
"type"
]
),
self.config["LLM"][self.config["selected_module"]["LLM"]], self.config["LLM"][self.config["selected_module"]["LLM"]],
), ),
tts.create_instance( tts.create_instance(
self.config["selected_module"]["TTS"] (
if not 'type' in self.config["TTS"][self.config["selected_module"]["TTS"]] self.config["selected_module"]["TTS"]
else if not "type"
self.config["TTS"][self.config["selected_module"]["TTS"]]["type"], in self.config["TTS"][self.config["selected_module"]["TTS"]]
else self.config["TTS"][self.config["selected_module"]["TTS"]][
"type"
]
),
self.config["TTS"][self.config["selected_module"]["TTS"]], self.config["TTS"][self.config["selected_module"]["TTS"]],
self.config["delete_audio"] self.config["delete_audio"],
), ),
memory.create_instance(memory_cls_name, memory_cfg), memory.create_instance(memory_cls_name, memory_cfg),
intent.create_instance( intent.create_instance(
self.config["selected_module"]["Intent"] (
if not 'type' in self.config["Intent"][self.config["selected_module"]["Intent"]] self.config["selected_module"]["Intent"]
else if not "type"
self.config["Intent"][self.config["selected_module"]["Intent"]]["type"], in self.config["Intent"][self.config["selected_module"]["Intent"]]
self.config["Intent"][self.config["selected_module"]["Intent"]] else self.config["Intent"][
self.config["selected_module"]["Intent"]
]["type"]
),
self.config["Intent"][self.config["selected_module"]["Intent"]],
), ),
) )
@@ -64,19 +86,33 @@ class WebSocketServer:
host = server_config["ip"] host = server_config["ip"]
port = server_config["port"] port = server_config["port"]
self.logger.bind(tag=TAG).info("Server is running at ws://{}:{}", get_local_ip(), port) self.logger.bind(tag=TAG).info(
self.logger.bind(tag=TAG).info("=======上面的地址是websocket协议地址,请勿用浏览器访问=======") "Server is running at ws://{}:{}/xiaozhi/v1/", get_local_ip(), port
async with websockets.serve( )
self._handle_connection, self.logger.bind(tag=TAG).info(
host, "=======上面的地址是websocket协议地址,请勿用浏览器访问======="
port )
): self.logger.bind(tag=TAG).info(
"如想测试websocket请用谷歌浏览器打开test目录下的test_page.html"
)
self.logger.bind(tag=TAG).info(
"=============================================================\n"
)
async with websockets.serve(self._handle_connection, host, port):
await asyncio.Future() await asyncio.Future()
async def _handle_connection(self, websocket): async def _handle_connection(self, websocket):
"""处理新连接,每次创建独立的ConnectionHandler""" """处理新连接,每次创建独立的ConnectionHandler"""
# 创建ConnectionHandler时传入当前server实例 # 创建ConnectionHandler时传入当前server实例
handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._memory, self.intent) handler = ConnectionHandler(
self.config,
self._vad,
self._asr,
self._llm,
self._tts,
self._memory,
self.intent,
)
self.active_connections.add(handler) self.active_connections.add(handler)
try: try:
await handler.handle_connection(websocket) await handler.handle_connection(websocket)