refactor: introduce transport-neutral connection runtime

This commit is contained in:
caixypromise
2026-07-27 02:05:26 +08:00
parent f5ed1aaec8
commit 0c582ed3b6
56 changed files with 9931 additions and 237 deletions
@@ -109,7 +109,7 @@ class ASRProvider(ASRProviderBase):
self.token, expire_time_str = AccessToken.create_token(self.access_key_id, self.access_key_secret)
if not self.token:
raise ValueError("无法获取有效的访问Token")
try:
expire_str = str(expire_time_str).strip()
if expire_str.isdigit():
@@ -151,7 +151,7 @@ class ASRProvider(ASRProviderBase):
"""开始识别会话"""
if self._is_token_expired():
self._refresh_token()
# 建立连接
headers = {"X-NLS-Token": self.token}
self.asr_ws = await websockets.connect(
@@ -169,7 +169,10 @@ class ASRProvider(ASRProviderBase):
self.is_processing = True
self.server_ready = False # 重置服务器准备状态
self.forward_task = asyncio.create_task(self._forward_results(conn))
session_id = getattr(conn, "session_id", None)
self.forward_task = self._create_session_task(
conn, self._forward_results(conn, session_id)
)
# 发送开始请求
start_request = {
@@ -193,10 +196,13 @@ class ASRProvider(ASRProviderBase):
await self.asr_ws.send(json.dumps(start_request, ensure_ascii=False))
logger.bind(tag=TAG).debug("已发送开始请求,等待服务器准备...")
async def _forward_results(self, conn: "ConnectionHandler"):
async def _forward_results(self, conn: "ConnectionHandler", session_id=None):
"""转发识别结果"""
try:
while not conn.stop_event.is_set():
while (
not conn.stop_event.is_set()
and self._session_is_current(conn, session_id)
):
# 获取当前连接的音频数据
audio_data = conn.asr_audio
try:
@@ -276,7 +282,7 @@ class ASRProvider(ASRProviderBase):
finally:
# 清理连接的音频缓存
await self._cleanup()
conn.reset_audio_states()
self._reset_audio_if_current(conn, session_id)
async def _send_stop_request(self):
"""发送停止识别请求(不关闭连接)"""
@@ -308,6 +314,21 @@ class ASRProvider(ASRProviderBase):
self.server_ready = False
logger.bind(tag=TAG).debug("ASR状态已重置")
forward_task = self.forward_task
current_task = asyncio.current_task()
if (
forward_task
and forward_task is not current_task
and not forward_task.done()
):
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"等待ASR转发任务退出失败: {e}")
# 关闭连接
if self.asr_ws:
try:
@@ -319,8 +340,10 @@ class ASRProvider(ASRProviderBase):
finally:
self.asr_ws = None
# 清理任务引用
self.forward_task = None
# Never discard a live task reference. The current forward task reaches
# this branch from its own finally block and is already completing.
if self.forward_task is forward_task:
self.forward_task = None
logger.bind(tag=TAG).debug("ASR会话清理完成")
@@ -103,7 +103,10 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).debug("WebSocket连接建立成功")
self.server_ready = False
self.forward_task = asyncio.create_task(self._forward_results(conn))
session_id = getattr(conn, "session_id", None)
self.forward_task = self._create_session_task(
conn, self._forward_results(conn, session_id)
)
# 发送run-task指令
run_task_msg = self._build_run_task_message()
@@ -154,10 +157,13 @@ class ASRProvider(ASRProviderBase):
return message
async def _forward_results(self, conn: "ConnectionHandler"):
async def _forward_results(self, conn: "ConnectionHandler", session_id=None):
"""转发识别结果"""
try:
while not conn.stop_event.is_set():
while (
not conn.stop_event.is_set()
and self._session_is_current(conn, session_id)
):
# 获取当前连接的音频数据
audio_data = conn.asr_audio
try:
@@ -243,7 +249,7 @@ class ASRProvider(ASRProviderBase):
finally:
# 清理连接的音频缓存
await self._cleanup()
conn.reset_audio_states()
self._reset_audio_if_current(conn, session_id)
async def _send_stop_request(self):
"""发送停止请求(用于手动模式停止录音)"""
@@ -285,6 +291,21 @@ class ASRProvider(ASRProviderBase):
self.server_ready = False
logger.bind(tag=TAG).debug("ASR状态已重置")
forward_task = self.forward_task
current_task = asyncio.current_task()
if (
forward_task
and forward_task is not current_task
and not forward_task.done()
):
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"等待ASR转发任务退出失败: {e}")
# 关闭连接
if self.asr_ws:
try:
@@ -301,8 +322,8 @@ class ASRProvider(ASRProviderBase):
finally:
self.asr_ws = None
# 清理任务引用
self.forward_task = None
if self.forward_task is forward_task:
self.forward_task = None
self.task_id = None
logger.bind(tag=TAG).debug("ASR会话清理完成")
@@ -315,4 +336,4 @@ class ASRProvider(ASRProviderBase):
async def close(self):
"""关闭资源"""
await self._cleanup()
await self._cleanup()
+54 -9
View File
@@ -32,8 +32,37 @@ class ASRProviderBase(ABC):
def __init__(self):
pass
@staticmethod
def _session_is_current(conn: "ConnectionHandler", session_id: str) -> bool:
"""Return whether an async ASR result still belongs to this session."""
return session_id is None or getattr(conn, "session_id", None) == session_id
def _create_session_task(self, conn: "ConnectionHandler", coroutine):
"""Bind streaming ASR work to the logical session when supported."""
creator = getattr(conn, "create_background_task", None)
if callable(creator):
try:
return creator(coroutine, turn_scoped=True)
except TypeError:
return creator(coroutine)
return asyncio.create_task(coroutine)
def _reset_audio_if_current(
self, conn: "ConnectionHandler", session_id: str
) -> None:
if not self._session_is_current(conn, session_id):
return
reset_audio_states = getattr(conn, "reset_audio_states", None)
if callable(reset_audio_states):
reset_audio_states()
# 打开音频通道
async def open_audio_channels(self, conn: "ConnectionHandler"):
# The pipeline runtime feeds PCM directly and does not use the legacy
# priority queue. Starting that worker would route frames back through
# ConnectionHandler-only functions.
if getattr(conn, "uses_pipeline_runtime", False):
return
conn.asr_priority_thread = threading.Thread(
target=self.asr_text_priority_thread, args=(conn,), daemon=True
)
@@ -72,18 +101,26 @@ class ASRProviderBase(ABC):
return
# 自动模式下通过VAD检测到语音停止时触发识别
if conn.asr.interface_type != InterfaceType.STREAM and conn.client_voice_stop:
interface_type = getattr(
self,
"interface_type",
getattr(getattr(conn, "asr", None), "interface_type", None),
)
if interface_type != InterfaceType.STREAM and conn.client_voice_stop:
# 直接使用asr_audio中的PCM数据
pcm_bytes = b"".join(conn.asr_audio)
# 检查是否有足够的音频数据(每帧1920字节,15帧约28800字节)
if len(pcm_bytes) > 1920 * 15:
await self.handle_voice_stop(conn, [pcm_bytes])
conn.reset_audio_states()
reset_audio_states = getattr(conn, "reset_audio_states", None)
if callable(reset_audio_states):
reset_audio_states()
# 处理语音停止
async def handle_voice_stop(self, conn: "ConnectionHandler", asr_audio_task: List[bytes]):
"""并行处理ASR和声纹识别"""
try:
session_id = getattr(conn, "session_id", None)
total_start_time = time.monotonic()
# 数据已经是PCM直接使用
@@ -96,13 +133,11 @@ class ASRProviderBase(ABC):
wav_data = self._pcm_to_wav(combined_pcm_data)
# 定义ASR任务
asr_task = self.speech_to_text_wrapper(
asr_audio_task, conn.session_id
)
asr_task = self.speech_to_text_wrapper(asr_audio_task, session_id)
if conn.voiceprint_provider and wav_data:
voiceprint_task = conn.voiceprint_provider.identify_speaker(
wav_data, conn.session_id
wav_data, session_id
)
# 并发等待两个结果
asr_result, voiceprint_result = await asyncio.gather(
@@ -112,6 +147,12 @@ class ASRProviderBase(ABC):
asr_result = await asr_task
voiceprint_result = None
if not self._session_is_current(conn, session_id):
logger.bind(tag=TAG).info(
f"丢弃旧会话ASR结果: session_id={session_id}"
)
return
# 记录识别结果 - 检查是否为异常
if isinstance(asr_result, Exception):
logger.bind(tag=TAG).error(f"ASR识别失败: {asr_result}")
@@ -165,9 +206,13 @@ class ASRProviderBase(ABC):
if text_len > 0:
audio_snapshot = asr_audio_task.copy()
enqueue_asr_report(conn, enhanced_text, audio_snapshot)
# 使用自定义模块进行上报
await startToChat(conn, enhanced_text)
result_handler = getattr(conn, "asr_result_handler", None)
if callable(result_handler):
await result_handler(enhanced_text, audio_snapshot)
else:
enqueue_asr_report(conn, enhanced_text, audio_snapshot)
# Legacy ConnectionHandler path.
await startToChat(conn, enhanced_text)
except Exception as e:
logger.bind(tag=TAG).error(f"处理语音停止失败: {e}")
import traceback
@@ -117,7 +117,10 @@ class ASRProvider(ASRProviderBase):
raise e
# 启动接收ASR结果的异步任务
self.forward_task = asyncio.create_task(self._forward_asr_results(conn))
session_id = getattr(conn, "session_id", None)
self.forward_task = self._create_session_task(
conn, self._forward_asr_results(conn, session_id)
)
# 发送缓存的音频数据
if conn.asr_audio and len(conn.asr_audio) > 0:
@@ -156,9 +159,13 @@ class ASRProvider(ASRProviderBase):
except Exception as e:
logger.bind(tag=TAG).info(f"发送音频数据时发生错误: {e}")
async def _forward_asr_results(self, conn: "ConnectionHandler"):
async def _forward_asr_results(self, conn: "ConnectionHandler", session_id=None):
try:
while self.asr_ws and not conn.stop_event.is_set():
while (
self.asr_ws
and not conn.stop_event.is_set()
and self._session_is_current(conn, session_id)
):
# 获取当前连接的音频数据
audio_data = conn.asr_audio
try:
@@ -249,21 +256,47 @@ class ASRProvider(ASRProviderBase):
if hasattr(e, "__cause__") and e.__cause__:
logger.bind(tag=TAG).error(f"错误原因: {str(e.__cause__)}")
finally:
if self.asr_ws:
await self.asr_ws.close()
self.asr_ws = None
self.is_processing = False
self._is_stopping = False
await self._cleanup()
# 重置所有音频相关状态
conn.reset_audio_states()
self._reset_audio_if_current(conn, session_id)
def stop_ws_connection(self):
if self.asr_ws:
asyncio.create_task(self.asr_ws.close())
self.asr_ws = None
# The forward task owns the WebSocket and closes it from _cleanup().
# Scheduling an untracked close here races with that cleanup path.
self.is_processing = False
self._is_stopping = False
async def _cleanup(self):
"""取消转发任务并关闭流式 ASR 连接。"""
self.is_processing = False
self._is_stopping = False
forward_task = self.forward_task
current_task = asyncio.current_task()
if (
forward_task
and forward_task is not current_task
and not forward_task.done()
):
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"等待ASR转发任务退出失败: {e}")
if self.asr_ws:
try:
await asyncio.wait_for(self.asr_ws.close(), timeout=2.0)
except Exception as e:
logger.bind(tag=TAG).warning(f"关闭ASR WebSocket连接失败: {e}")
finally:
self.asr_ws = None
if self.forward_task is forward_task:
self.forward_task = None
async def _send_stop_request(self):
"""发送最后一个音频帧以通知服务器结束"""
self._is_stopping = True # 先标记为停止状态,阻止后续音频发送
@@ -417,14 +450,4 @@ class ASRProvider(ASRProviderBase):
async def close(self):
"""资源清理方法"""
if self.asr_ws:
await self.asr_ws.close()
self.asr_ws = None
if self.forward_task:
self.forward_task.cancel()
try:
await self.forward_task
except asyncio.CancelledError:
pass
self.forward_task = None
self.is_processing = False
await self._cleanup()
@@ -90,8 +90,11 @@ class ASRProvider(ASRProviderBase):
batch_size_s=60,
)
text = lang_tag_filter(result[0]["text"])
recognized_content = (
text.get("content", "") if isinstance(text, dict) else text
)
logger.bind(tag=TAG).debug(
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text['content']}"
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {recognized_content}"
)
return text, artifacts.file_path
@@ -0,0 +1,522 @@
"""
SharedASRManager: 全局 ASR 管理器
实现单例模型 + 单推理执行器 + 队列限流。
单例的原因是:推理是 CPU/GPU-bound,不是 I/O-bound,多实例不仅会占用内存,还会降低吞吐能力
"""
import asyncio
import copy
import json
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, Any, Optional, Tuple, List
from config.logger import setup_logging
logger = setup_logging()
TAG = __name__
class SharedASRManager:
"""
全局共享 ASR 管理器
"""
# 支持预加载的本地模型类型
LOCAL_MODEL_TYPES = [
"fun_local", # FunASR 本地
"sherpa_onnx_local", # Sherpa ONNX
"sense_voice" # SenseVoice
]
def __init__(self, config: Dict[str, Any], asr_type: str = None):
"""
初始化 ASR 管理器
Args:
config: 服务器配置
asr_type: ASR 类型(Optional,用于显式指定)
"""
self.config = config
self.asr_type = asr_type
# 模型实例(全局单例)
self.model_instance = None
# 任务队列(限流)
queue_max_size = self._get_queue_max_size()
self.task_queue: asyncio.Queue = asyncio.Queue(maxsize=queue_max_size)
# 推理锁(使得推理串行化)
self.inference_lock = asyncio.Lock()
# 线程池执行器,用于阻塞调用
self.executor: Optional[ThreadPoolExecutor] = None
# 运行状态
self.running = False
self._inference_task: Optional[asyncio.Task] = None
self.is_local_model = self._check_local_model()
self._variant_lock = asyncio.Lock()
self._variant_managers = {}
self._max_shared_models = max(
1, int(config.get("shared_asr_max_models", 3) or 3)
)
logger.bind(tag=TAG).info(
f"SharedASRManager 初始化完成, "
f"类型: {self.asr_type}, "
f"本地模型: {self.is_local_model}, "
f"队列大小: {queue_max_size}"
)
def _get_queue_max_size(self) -> int:
"""获取队列最大大小"""
# 尝试从配置获取
selected_asr = self.config.get("selected_module", {}).get("ASR")
if selected_asr:
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
return asr_config.get("queue_max_size", 100)
return 100
def _check_local_model(self) -> bool:
"""检查是否为本地模型"""
if self.asr_type:
return self.asr_type in self.LOCAL_MODEL_TYPES
# 从配置推断
selected_asr = self.config.get("selected_module", {}).get("ASR")
if not selected_asr:
return False
asr_config = self.config.get("ASR", {}).get(selected_asr, {})
asr_type = asr_config.get("type", selected_asr)
self.asr_type = asr_type
return asr_type in self.LOCAL_MODEL_TYPES
async def initialize(self):
"""
初始化管理器
- 预加载模型
- 启动推理执行器
"""
if not self.is_local_model:
logger.bind(tag=TAG).info("非本地模型,跳过预加载")
return
if self.running:
logger.bind(tag=TAG).warning("管理器已在运行中")
return
try:
logger.bind(tag=TAG).info(f"开始预加载 ASR 模型: {self.asr_type}")
# 预加载模型
await self._preload_model()
# 启动推理执行器
self.running = True
self._inference_task = asyncio.create_task(self._inference_loop())
logger.bind(tag=TAG).info("ASR 模型预加载完成,推理执行器已启动")
except Exception as e:
logger.bind(tag=TAG).error(f"ASR 模型预加载失败: {e}")
raise
async def _preload_model(self):
"""在线程池中预加载模型"""
# 创建线程池
self.executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="asr_worker")
loop = asyncio.get_event_loop()
self.model_instance = await loop.run_in_executor(
self.executor,
self._create_model_instance
)
logger.bind(tag=TAG).info("模型实例创建完成")
def _create_model_instance(self):
"""
实际创建模型实例(在线程中执行)
Returns:
ASR Provider 实例
"""
from core.utils.modules_initialize import initialize_asr
logger.bind(tag=TAG).info("正在创建 ASR 模型实例...")
instance = initialize_asr(self.config)
logger.bind(tag=TAG).info("ASR 模型实例创建成功")
return instance
async def submit_task(
self,
opus_data: List[bytes],
session_id: str,
audio_format: str = "opus"
) -> Tuple[Optional[str], Optional[str]]:
"""
提交推理任务
Args:
opus_data: 音频数据
session_id: 会话 ID
audio_format: 音频格式
Returns:
(识别文本, 文件路径)
Raises:
RuntimeError: 队列满或服务未运行
"""
if not self.running:
raise RuntimeError("ASR 服务未运行")
# 检查队列是否满(限流)
if self.task_queue.full():
queue_status = self.get_queue_status()
logger.bind(tag=TAG).warning(
f"ASR 队列已满: {queue_status}"
)
raise RuntimeError("ASR 服务繁忙,请稍后重试")
# 创建 Future 用于返回结果
result_future: asyncio.Future = asyncio.Future()
# 构造任务
task = {
'opus_data': opus_data,
'session_id': session_id,
'audio_format': audio_format,
'future': result_future
}
# 放入队列
await self.task_queue.put(task)
logger.bind(tag=TAG).debug(
f"任务已提交, session: {session_id}, "
f"队列大小: {self.task_queue.qsize()}"
)
# 等待结果
return await result_future
async def _inference_loop(self):
"""
单个推理执行器循环
核心原则:
- 只有一个执行器
- 串行处理任务
- 带超时的队列获取
"""
logger.bind(tag=TAG).info("推理执行器启动")
while self.running:
task = None
try:
# 带超时的队列获取,避免关闭时卡住
try:
task = await asyncio.wait_for(
self.task_queue.get(),
timeout=1.0
)
except asyncio.TimeoutError:
# 超时后检查 running 状态,继续循环
continue
if task['future'].done():
continue
# 执行推理(加锁保证串行)
async with self.inference_lock:
if task['future'].done():
continue
result = await self._run_inference(
task['opus_data'],
task['session_id'],
task['audio_format']
)
# 返回结果
if not task['future'].done():
task['future'].set_result(result)
logger.bind(tag=TAG).debug(
f"推理完成, session: {task['session_id']}"
)
except asyncio.CancelledError:
logger.bind(tag=TAG).info("推理执行器被取消")
break
except Exception as e:
logger.bind(tag=TAG).error(f"推理执行失败: {e}")
if task and 'future' in task and not task['future'].done():
task['future'].set_exception(e)
finally:
if task is not None:
self.task_queue.task_done()
logger.bind(tag=TAG).info("推理执行器已停止")
async def _run_inference(
self,
opus_data: List[bytes],
session_id: str,
audio_format: str
) -> Tuple[Optional[str], Optional[str]]:
"""
执行实际推理(在线程池中)
Args:
opus_data: 音频数据
session_id: 会话 ID
audio_format: 音频格式
Returns:
(识别文本, 文件路径)
"""
loop = asyncio.get_event_loop()
# 在线程池中执行推理
result = await loop.run_in_executor(
self.executor,
lambda: self._sync_wrapper(opus_data, session_id)
)
return result
def _sync_wrapper(
self,
opus_data: List[bytes],
session_id: str,
) -> Tuple[Optional[str], Optional[str]]:
"""Run the provider's complete artifact wrapper in the worker thread."""
import asyncio
async def _call():
return await self.model_instance.speech_to_text_wrapper(
opus_data, session_id
)
# 创建新的事件循环执行
loop = None
try:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
result = loop.run_until_complete(_call())
return result
finally:
if loop:
loop.close()
def matches_config(self, config: Dict[str, Any]) -> bool:
"""Return whether this manager owns the exact selected ASR config."""
selected = config.get("selected_module", {}).get("ASR")
manager_selected = self.config.get("selected_module", {}).get("ASR")
if not selected or selected != manager_selected:
return False
return (
config.get("ASR", {}).get(selected)
== self.config.get("ASR", {}).get(manager_selected)
and bool(config.get("delete_audio", True))
== bool(self.config.get("delete_audio", True))
)
@classmethod
def _config_fingerprint(cls, config: Dict[str, Any]) -> str:
selected = config.get("selected_module", {}).get("ASR")
payload = {
"selected": selected,
"config": config.get("ASR", {}).get(selected),
"delete_audio": bool(config.get("delete_audio", True)),
}
return json.dumps(payload, sort_keys=True, separators=(",", ":"))
@classmethod
def _config_is_local(cls, config: Dict[str, Any]) -> bool:
selected = config.get("selected_module", {}).get("ASR")
asr_config = config.get("ASR", {}).get(selected, {})
asr_type = asr_config.get("type", selected)
return asr_type in cls.LOCAL_MODEL_TYPES
async def acquire_for_config(
self, config: Dict[str, Any]
) -> Optional["SharedASRManager"]:
"""Acquire a bounded shared local model matching one Agent config."""
if self.matches_config(config):
return self if self.is_ready() else None
if not self._config_is_local(config):
return None
fingerprint = self._config_fingerprint(config)
async with self._variant_lock:
entry = self._variant_managers.get(fingerprint)
if entry is not None:
entry["references"] += 1
return entry["manager"]
if len(self._variant_managers) >= self._max_shared_models - 1:
idle_fingerprint = next(
(
key
for key, candidate in self._variant_managers.items()
if candidate["references"] == 0
),
None,
)
if idle_fingerprint is None:
raise RuntimeError(
"ASR共享模型容量已满,请提高shared_asr_max_models或统一Agent ASR配置"
)
idle_manager = self._variant_managers.pop(
idle_fingerprint
)["manager"]
# Keep the lock until shutdown completes so a replacement can
# never overlap the retiring model and exceed the hard limit.
await idle_manager.shutdown()
manager = SharedASRManager(copy.deepcopy(config))
try:
await manager.initialize()
except Exception:
await manager.shutdown()
raise
self._variant_managers[fingerprint] = {
"manager": manager,
"references": 1,
}
logger.bind(tag=TAG).info(
"已加载Agent专用共享ASR模型,当前模型数: {}",
len(self._variant_managers) + 1,
)
return manager
async def release_for_config(self, manager: "SharedASRManager") -> None:
"""Release an Agent model; idle variants stay cached for safe reuse."""
if manager is self:
return
async with self._variant_lock:
for entry in self._variant_managers.values():
if entry["manager"] is not manager:
continue
entry["references"] = max(0, entry["references"] - 1)
break
async def shutdown(self):
"""
优雅停机
步骤:
1. 停止接收新任务
2. 等待当前任务完成(带超时)
3. 取消未完成的任务
4. 关闭线程池
"""
async with self._variant_lock:
variants = [
entry["manager"] for entry in self._variant_managers.values()
]
self._variant_managers.clear()
if variants:
await asyncio.gather(
*(manager.shutdown() for manager in variants),
return_exceptions=True,
)
if (
not self.running
and self.executor is None
and self.model_instance is None
):
return
logger.bind(tag=TAG).info("开始关闭 ASR 管理器...")
# 停止接收新任务
self.running = False
# 等待推理任务完成
if self._inference_task and not self._inference_task.done():
try:
# 最多等待 5 秒
await asyncio.wait_for(
self._inference_task,
timeout=5.0
)
except asyncio.TimeoutError:
logger.bind(tag=TAG).warning("推理任务超时,强制取消")
self._inference_task.cancel()
try:
await self._inference_task
except asyncio.CancelledError:
pass
except asyncio.CancelledError:
pass
# 取消所有队列中未完成的任务
cancelled_count = 0
while not self.task_queue.empty():
try:
task = self.task_queue.get_nowait()
if not task['future'].done():
task['future'].set_exception(
RuntimeError("ASR 服务正在关闭")
)
cancelled_count += 1
except asyncio.QueueEmpty:
break
if cancelled_count > 0:
logger.bind(tag=TAG).info(f"已取消 {cancelled_count} 个待处理任务")
# 关闭线程池
if self.executor:
self.executor.shutdown(wait=False)
self.executor = None
logger.bind(tag=TAG).info("线程池已关闭")
# 清理模型实例
self.model_instance = None
logger.bind(tag=TAG).info("ASR 管理器已关闭")
def get_queue_status(self) -> Dict[str, Any]:
"""
获取队列状态(用于监控)
Returns:
队列状态字典
"""
queue_size = self.task_queue.qsize()
queue_max = self.task_queue.maxsize
return {
'queue_size': queue_size,
'queue_max': queue_max,
'is_busy': queue_size > queue_max * 0.8,
'utilization': queue_size / queue_max if queue_max > 0 else 0,
'running': self.running
}
def is_ready(self) -> bool:
"""检查管理器是否就绪"""
return (
self.running and
self.model_instance is not None and
self.executor is not None
)
@classmethod
def is_local_model_type(cls, asr_type: str) -> bool:
"""
检查 ASR 类型是否为本地模型
Args:
asr_type: ASR 类型
Returns:
是否为本地模型
"""
return asr_type in cls.LOCAL_MODEL_TYPES
@@ -0,0 +1,119 @@
"""
SharedASRProxy: 共享 ASR 管理器的代理类
功能:
- 包装 SharedASRManager
- 提供与原 ASR Provider 相同的接口
- 处理队列满等异常情况
"""
from typing import List, Tuple, Optional, Dict, Any
from core.providers.asr.base import ASRProviderBase
from core.providers.asr.dto.dto import InterfaceType
from config.logger import setup_logging
logger = setup_logging()
TAG = __name__
class SharedASRProxy(ASRProviderBase):
"""
共享 ASR 管理器的代理类
该类提供与原 ASR Provider 相同的接口,
但实际推理工作由 SharedASRManager 完成。
"""
def __init__(self, manager):
"""
初始化代理
Args:
manager: SharedASRManager 实例
"""
super().__init__()
self.manager = manager
# 从共享管理器获取接口类型
if manager.model_instance and hasattr(manager.model_instance, 'interface_type'):
self.interface_type = manager.model_instance.interface_type
else:
self.interface_type = InterfaceType.LOCAL
logger.bind(tag=TAG).info("SharedASRProxy 初始化完成")
async def speech_to_text(
self,
opus_data: List[bytes],
session_id: str,
audio_format: str = "opus"
) -> Tuple[Optional[str], Optional[str]]:
"""
语音转文本(通过共享管理器)
Args:
opus_data: 音频数据(Opus 编码的字节列表)
session_id: 会话 ID
audio_format: 音频格式,默认 "opus"
Returns:
Tuple[str, str]: (识别的文本, 音频文件路径)
"""
try:
# 检查管理器状态
if not self.manager.is_ready():
logger.bind(tag=TAG).error("ASR 管理器未就绪")
return "", None
# 提交任务到共享管理器
result = await self.manager.submit_task(
opus_data,
session_id,
audio_format
)
return result
except RuntimeError as e:
# 队列满或服务未运行
logger.bind(tag=TAG).warning(f"ASR 服务繁忙: {e}")
# 返回友好提示,而不是空字符串
return "服务繁忙,请稍后重试", None
except Exception as e:
logger.bind(tag=TAG).error(f"ASR 推理失败: {e}")
return "", None
async def speech_to_text_wrapper(
self, pcm_data: List[bytes], session_id: str
) -> Tuple[Optional[str], Optional[str]]:
"""Delegate PCM inference without using uninitialized proxy file state."""
return await self.manager.submit_task(pcm_data, session_id, "pcm")
def get_queue_status(self) -> Dict[str, Any]:
"""
获取队列状态
Returns:
队列状态字典
"""
return self.manager.get_queue_status()
def is_ready(self) -> bool:
"""
检查代理是否就绪
Returns:
是否就绪
"""
return self.manager.is_ready()
async def close(self):
"""
关闭代理
注意:不关闭共享管理器,由服务器统一管理
"""
logger.bind(tag=TAG).debug("SharedASRProxy 关闭")
# 代理不负责关闭共享管理器
pass
@@ -141,7 +141,10 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).info("ASR WebSocket连接已建立")
self.server_ready = False
self.forward_task = asyncio.create_task(self._forward_results(conn))
session_id = getattr(conn, "session_id", None)
self.forward_task = self._create_session_task(
conn, self._forward_results(conn, session_id)
)
# 发送首帧音频
if conn.asr_audio and len(conn.asr_audio) > 0:
@@ -185,10 +188,13 @@ class ASRProvider(ASRProviderBase):
await self.asr_ws.send(json.dumps(frame_data, ensure_ascii=False))
async def _forward_results(self, conn: "ConnectionHandler"):
async def _forward_results(self, conn: "ConnectionHandler", session_id=None):
"""转发识别结果"""
try:
while not conn.stop_event.is_set():
while (
not conn.stop_event.is_set()
and self._session_is_current(conn, session_id)
):
try:
response = await asyncio.wait_for(self.asr_ws.recv(), timeout=60)
result = json.loads(response)
@@ -247,7 +253,7 @@ class ASRProvider(ASRProviderBase):
finally:
# 清理连接资源
await self._cleanup()
conn.reset_audio_states()
self._reset_audio_if_current(conn, session_id)
async def handle_voice_stop(
self, conn: "ConnectionHandler", asr_audio_task: List[bytes]
@@ -272,9 +278,8 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).debug(f"异常详情: {traceback.format_exc()}")
def stop_ws_connection(self):
if self.asr_ws:
asyncio.create_task(self.asr_ws.close())
self.asr_ws = None
# The forward task owns the WebSocket and closes it from _cleanup().
# Scheduling an untracked close here races with that cleanup path.
self.is_processing = False
async def _send_stop_request(self):
@@ -299,6 +304,21 @@ class ASRProvider(ASRProviderBase):
self.server_ready = False
logger.bind(tag=TAG).debug("ASR状态已重置")
forward_task = self.forward_task
current_task = asyncio.current_task()
if (
forward_task
and forward_task is not current_task
and not forward_task.done()
):
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"等待ASR转发任务退出失败: {e}")
# 关闭连接
if self.asr_ws:
try:
@@ -310,8 +330,8 @@ class ASRProvider(ASRProviderBase):
finally:
self.asr_ws = None
# 清理任务引用
self.forward_task = None
if self.forward_task is forward_task:
self.forward_task = None
logger.bind(tag=TAG).debug("ASR会话清理完成")
@@ -323,15 +343,4 @@ class ASRProvider(ASRProviderBase):
async def close(self):
"""资源清理方法"""
if self.asr_ws:
await self.asr_ws.close()
self.asr_ws = None
if self.forward_task:
self.forward_task.cancel()
try:
await self.forward_task
except asyncio.CancelledError:
pass
self.forward_task = None
self.is_processing = False
await self._cleanup()
@@ -127,7 +127,16 @@ class DeviceIoTExecutor(ToolExecutor):
send_message = json.dumps(
{"type": "iot", "commands": [command]}
)
await self.conn.websocket.send(send_message)
# 使用transport接口发送消息
if hasattr(self.conn, 'transport') and self.conn.transport:
await self.conn.transport.send(send_message)
elif hasattr(self.conn, 'websocket') and self.conn.websocket:
# 兼容旧版本
logger.warning("未找到SessionContext的传输层接口, 回退使用旧版conn.websocket发送消息")
await self.conn.websocket.send(send_message)
else:
raise AttributeError("无法找到可用的传输层接口")
return
raise Exception(f"未找到设备{device_name}的方法{method_name}")
@@ -17,7 +17,7 @@ class MCPClient:
self.name_mapping = {}
self.ready = False
self.call_results = {} # To store Futures for tool call responses
self.next_id = 1
self.next_id = 10000
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
@@ -91,3 +91,12 @@ class MCPClient:
async with self.lock:
if id in self.call_results:
self.call_results.pop(id)
async def close(self):
async with self.lock:
pending = list(self.call_results.values())
self.call_results.clear()
self.ready = False
for future in pending:
if not future.done():
future.set_exception(ConnectionError("MCP会话已关闭"))
@@ -3,11 +3,11 @@
import json
import asyncio
import re
from concurrent.futures import Future
from core.utils.util import get_vision_url, sanitize_tool_name
from core.utils.util import get_vision_url
from core.utils.auth import AuthToken
from config.logger import setup_logging
from typing import TYPE_CHECKING
from .mcp_client import MCPClient
if TYPE_CHECKING:
from core.connection import ConnectionHandler
@@ -15,108 +15,45 @@ if TYPE_CHECKING:
TAG = __name__
logger = setup_logging()
class MCPClient:
"""设备端MCP客户端,用于管理MCP状态和工具"""
def __init__(self):
self.tools = {} # sanitized_name -> tool_data
self.name_mapping = {}
self.ready = False
self.call_results = {} # To store Futures for tool call responses
self.next_id = 1
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
def has_tool(self, name: str) -> bool:
return name in self.tools
def get_available_tools(self) -> list:
# Check if the cache is valid
if self._cached_available_tools is not None:
return self._cached_available_tools
# If cache is not valid, regenerate the list
result = []
for tool_name, tool_data in self.tools.items():
function_def = {
"name": tool_name,
"description": tool_data["description"],
"parameters": {
"type": tool_data["inputSchema"].get("type", "object"),
"properties": tool_data["inputSchema"].get("properties", {}),
"required": tool_data["inputSchema"].get("required", []),
},
}
result.append({"type": "function", "function": function_def})
self._cached_available_tools = result # Store the generated list in cache
return result
async def is_ready(self) -> bool:
async with self.lock:
return self.ready
async def set_ready(self, status: bool):
async with self.lock:
self.ready = status
async def add_tool(self, tool_data: dict):
async with self.lock:
sanitized_name = sanitize_tool_name(tool_data["name"])
self.tools[sanitized_name] = tool_data
self.name_mapping[sanitized_name] = tool_data["name"]
self._cached_available_tools = (
None # Invalidate the cache when a tool is added
)
async def get_next_id(self) -> int:
async with self.lock:
current_id = self.next_id
self.next_id += 1
return current_id
async def register_call_result_future(self, id: int, future: Future):
async with self.lock:
self.call_results[id] = future
async def resolve_call_result(self, id: int, result: any):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_result(result)
async def reject_call_result(self, id: int, exception: Exception):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_exception(exception)
async def cleanup_call_result(self, id: int):
async with self.lock:
if id in self.call_results:
self.call_results.pop(id)
async def send_mcp_message(conn: "ConnectionHandler", payload: dict):
async def send_mcp_message(
conn: "ConnectionHandler",
payload: dict,
transport=None,
*,
raise_on_error: bool = False,
):
"""Helper to send MCP messages, encapsulating common logic."""
if not conn.features.get("mcp"):
features = getattr(conn, "features", {}) or {}
if not features.get("mcp"):
logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
return
message = json.dumps({"type": "mcp", "payload": payload})
try:
await conn.websocket.send(message)
# 优先使用传入的transport,否则尝试从conn获取
if transport:
await transport.send(message)
elif getattr(conn, "transport", None):
# 新架构
await conn.transport.send(message)
elif getattr(conn, "websocket", None):
# 兼容旧版本
await conn.websocket.send(message)
else:
raise AttributeError("无法找到可用的传输层接口")
logger.bind(tag=TAG).debug(f"成功发送MCP消息: {message}")
except Exception as e:
logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}")
if raise_on_error:
raise
async def handle_mcp_message(
conn: "ConnectionHandler", mcp_client: MCPClient, payload: dict
conn: "ConnectionHandler",
mcp_client: MCPClient,
payload: dict,
transport=None,
):
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
logger.bind(tag=TAG).debug(f"处理MCP消息: {str(payload)[:100]}")
@@ -150,7 +87,7 @@ async def handle_mcp_message(
await asyncio.sleep(1)
logger.bind(tag=TAG).debug("初始化完成,开始请求MCP工具列表")
await send_mcp_tools_list_request(conn)
await send_mcp_tools_list_request(conn, transport)
return
@@ -207,7 +144,7 @@ async def handle_mcp_message(
next_cursor = result.get("nextCursor", "")
if next_cursor:
logger.bind(tag=TAG).debug(f"有更多工具,nextCursor: {next_cursor}")
await send_mcp_tools_list_continue_request(conn, next_cursor)
await send_mcp_tools_list_continue_request(conn, next_cursor, transport)
else:
await mcp_client.set_ready(True)
logger.bind(tag=TAG).debug("所有工具已获取,MCP客户端准备就绪")
@@ -235,7 +172,9 @@ async def handle_mcp_message(
)
async def send_mcp_initialize_message(conn: "ConnectionHandler"):
async def send_mcp_initialize_message(
conn: "ConnectionHandler", transport=None
):
"""发送MCP初始化消息"""
vision_url = get_vision_url(conn.config)
@@ -267,10 +206,12 @@ async def send_mcp_initialize_message(conn: "ConnectionHandler"):
},
}
logger.bind(tag=TAG).debug("发送MCP初始化消息")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def send_mcp_tools_list_request(conn: "ConnectionHandler"):
async def send_mcp_tools_list_request(
conn: "ConnectionHandler", transport=None
):
"""发送MCP工具列表请求"""
payload = {
"jsonrpc": "2.0",
@@ -278,10 +219,12 @@ async def send_mcp_tools_list_request(conn: "ConnectionHandler"):
"method": "tools/list",
}
logger.bind(tag=TAG).debug("发送MCP工具列表请求")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def send_mcp_tools_list_continue_request(conn: "ConnectionHandler", cursor: str):
async def send_mcp_tools_list_continue_request(
conn: "ConnectionHandler", cursor: str, transport=None
):
"""发送带有cursor的MCP工具列表请求"""
payload = {
"jsonrpc": "2.0",
@@ -290,7 +233,7 @@ async def send_mcp_tools_list_continue_request(conn: "ConnectionHandler", cursor
"params": {"cursor": cursor},
}
logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def call_mcp_tool(
@@ -299,6 +242,7 @@ async def call_mcp_tool(
tool_name: str,
args: str = "{}",
timeout: int = 30,
return_raw: bool = False,
):
"""
调用指定的工具,并等待响应
@@ -359,6 +303,7 @@ async def call_mcp_tool(
raise ValueError(f"参数必须是字典类型,实际类型: {type(arguments)}")
except Exception as e:
await mcp_client.cleanup_call_result(tool_call_id)
if not isinstance(e, ValueError):
raise ValueError(f"参数处理失败: {str(e)}")
raise e
@@ -372,7 +317,11 @@ async def call_mcp_tool(
}
logger.bind(tag=TAG).info(f"发送客户端mcp工具调用请求: {actual_name},参数: {args}")
await send_mcp_message(conn, payload)
try:
await send_mcp_message(conn, payload, raise_on_error=True)
except Exception:
await mcp_client.cleanup_call_result(tool_call_id)
raise
try:
# Wait for response or timeout
@@ -388,13 +337,21 @@ async def call_mcp_tool(
)
raise RuntimeError(f"工具调用错误: {error_msg}")
if return_raw:
return raw_result
content = raw_result.get("content")
if isinstance(content, list) and len(content) > 0:
if isinstance(content[0], dict) and "text" in content[0]:
# 直接返回文本内容,不进行JSON解析
return content[0]["text"]
# 如果结果不是预期的格式,将其转换为字符串
if return_raw:
return raw_result
return str(raw_result)
except asyncio.CancelledError:
await mcp_client.cleanup_call_result(tool_call_id)
raise
except asyncio.TimeoutError:
await mcp_client.cleanup_call_result(tool_call_id)
raise TimeoutError("工具调用请求超时")
@@ -22,6 +22,7 @@ class MCPEndpointClient:
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
self.websocket = None # WebSocket连接
self.listener_task = None
def has_tool(self, name: str) -> bool:
return name in self.tools
@@ -107,6 +108,18 @@ class MCPEndpointClient:
async def close(self):
"""关闭WebSocket连接"""
current_task = asyncio.current_task()
if (
self.listener_task
and self.listener_task is not current_task
and not self.listener_task.done()
):
self.listener_task.cancel()
try:
await self.listener_task
except asyncio.CancelledError:
pass
self.listener_task = None
if self.websocket:
await self.websocket.close()
self.websocket = None
@@ -16,6 +16,8 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
if not mcp_endpoint_url or "你的" in mcp_endpoint_url or mcp_endpoint_url == "null":
return None
websocket = None
mcp_client = None
try:
websocket = await websockets.connect(mcp_endpoint_url)
@@ -23,7 +25,9 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
mcp_client.set_websocket(websocket)
# 启动消息监听器
asyncio.create_task(_message_listener(mcp_client))
mcp_client.listener_task = asyncio.create_task(
_message_listener(mcp_client), name="xiaozhi-mcp-endpoint-listener"
)
# 发送初始化消息
await send_mcp_endpoint_initialize(mcp_client)
@@ -39,6 +43,10 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
except Exception as e:
logger.bind(tag=TAG).error(f"连接MCP接入点失败: {e}")
if mcp_client is not None:
await mcp_client.close()
elif websocket is not None:
await websocket.close()
return None
@@ -72,16 +72,25 @@ class ServerMCPClient:
async def cleanup(self):
"""清理MCP客户端资源"""
if not self._worker_task:
task = self._worker_task
if not task:
return
self._shutdown_evt.set()
try:
await asyncio.wait_for(self._worker_task, timeout=20)
except (asyncio.TimeoutError, Exception) as e:
await asyncio.wait_for(asyncio.shield(task), timeout=20)
except asyncio.TimeoutError:
self.logger.bind(tag=TAG).warning("服务端MCP关闭超时,取消工作任务")
task.cancel()
done, _ = await asyncio.wait({task}, timeout=5)
if task not in done:
self.logger.bind(tag=TAG).error("服务端MCP工作任务取消超时")
return
except Exception as e:
self.logger.bind(tag=TAG).error(f"服务端MCP客户端关闭错误: {e}")
finally:
self._worker_task = None
if task.done():
self._worker_task = None
def has_tool(self, name: str) -> bool:
"""检查是否包含指定工具
@@ -197,7 +206,7 @@ class ServerMCPClient:
if "API_ACCESS_TOKEN" in self.config:
headers["Authorization"] = f"Bearer {self.config['API_ACCESS_TOKEN']}"
self.logger.bind(tag=TAG).warning(f"你正在使用旧过时的配置 API_ACCESS_TOKEN ,请在.mcp_server_settings.json中将API_ACCESS_TOKEN直接设置在headers中,例如 'Authorization': 'Bearer API_ACCESS_TOKEN'")
# 根据transport类型选择不同的客户端,默认为SSE
transport_type = self.config.get("transport", "sse")
+96 -7
View File
@@ -1,4 +1,5 @@
import os
import json
import re
import uuid
import queue
@@ -16,8 +17,6 @@ from config.logger import setup_logging
from core.utils import opus_encoder_utils
from core.utils.tts import MarkdownCleaner, convert_percentage_to_range
from core.utils.output_counter import add_device_output
from core.handle.reportHandle import enqueue_tts_report
from core.handle.sendAudioHandle import sendAudioMessage
from core.utils.util import audio_bytes_to_data_stream, audio_to_data_stream
from core.providers.tts.dto.dto import (
TTSMessageDTO,
@@ -30,6 +29,58 @@ TAG = __name__
logger = setup_logging()
async def sendAudioMessage(conn, sentenceType, audios, text, sentence_id=None):
"""兼容函数:使用新的processor发送音频消息"""
try:
# 获取transport接口
transport = getattr(conn, 'transport', None)
if not transport:
logger.error("SessionContext中没有transport接口")
return
# 使用AudioSendProcessor发送音频
from core.processors.audio_send_processor import AudioSendProcessor
processor = AudioSendProcessor()
await processor.send_audio_message(
conn,
transport,
sentenceType,
audios,
text,
sentence_id=sentence_id,
)
except Exception as e:
logger.error(f"发送音频消息失败: {e}")
import traceback
traceback.print_exc()
def enqueue_tts_report(conn, text, audio_data):
"""兼容函数:使用新的processor处理TTS报告"""
try:
# 获取transport接口
transport = getattr(conn, 'transport', None)
if not transport:
logger.error("SessionContext中没有transport接口")
return
# 使用ReportProcessor处理报告
from core.processors.report_processor import ReportProcessor
processor = ReportProcessor()
# 异步执行报告
if hasattr(conn, 'loop') and conn.loop:
# 直接调用同步方法
processor.enqueue_tts_report(conn, text, audio_data)
else:
logger.warning("SessionContext中没有事件循环,跳过TTS报告")
except Exception as e:
logger.error(f"TTS报告处理失败: {e}")
class TTSProviderBase(ABC):
def __init__(self, config, delete_audio_file):
self.interface_type = InterfaceType.NON_STREAM
@@ -189,7 +240,7 @@ class TTSProviderBase(ABC):
except Exception as e:
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
return None
def to_tts(self, text):
# 保留原始文本用于日志/显示
original_text = text
@@ -312,13 +363,17 @@ class TTSProviderBase(ABC):
# tts 消化线程
self.tts_priority_thread = threading.Thread(
target=self.tts_text_priority_thread, daemon=True
target=self.tts_text_priority_thread,
name=f"xiaozhi-tts-text-{id(self)}",
daemon=True,
)
self.tts_priority_thread.start()
# 音频播放 消化线程
self.audio_play_priority_thread = threading.Thread(
target=self._audio_play_priority_thread, daemon=True
target=self._audio_play_priority_thread,
name=f"xiaozhi-tts-audio-{id(self)}",
daemon=True,
)
self.audio_play_priority_thread.start()
@@ -369,6 +424,8 @@ class TTSProviderBase(ABC):
while not self.conn.stop_event.is_set():
try:
message = self.tts_text_queue.get(timeout=1)
if message is None:
break
if self.conn.client_abort:
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
continue
@@ -417,11 +474,15 @@ class TTSProviderBase(ABC):
try:
try:
item = self.tts_audio_queue.get(timeout=0.1)
if item is None:
break
if len(item) == 4:
sentence_type, audio_datas, text, sentence_id = item
else:
sentence_type, audio_datas, text = item
sentence_id = None
sentence_id = getattr(
self, "current_sentence_id", None
) or getattr(self.conn, "sentence_id", None)
except queue.Empty:
if self.conn.stop_event.is_set():
break
@@ -459,7 +520,19 @@ class TTSProviderBase(ABC):
sendAudioMessage(self.conn, sentence_type, audio_datas, text, sentence_id),
self.conn.loop,
)
future.result()
self._pending_audio_future = future
try:
future.result(timeout=max(1, self.tts_timeout))
except concurrent.futures.CancelledError:
break
except concurrent.futures.TimeoutError:
future.cancel()
logger.bind(tag=TAG).warning(
"TTS音频发送超时,取消当前发送任务"
)
finally:
if getattr(self, "_pending_audio_future", None) is future:
self._pending_audio_future = None
# 记录输出和报告
if self.conn.max_output_size > 0 and text:
@@ -476,6 +549,22 @@ class TTSProviderBase(ABC):
async def close(self):
"""资源清理方法"""
self.tts_stop_request = True
pending_future = getattr(self, "_pending_audio_future", None)
if pending_future and not pending_future.done():
pending_future.cancel()
# Wake workers immediately; relying on queue timeouts leaves Provider
# threads alive after the component has released its ownership.
self.tts_text_queue.put(None)
self.tts_audio_queue.put(None)
for thread_name in ("tts_priority_thread", "audio_play_priority_thread"):
thread = getattr(self, thread_name, None)
if thread and thread.is_alive() and thread is not threading.current_thread():
await asyncio.to_thread(thread.join, 2)
if thread.is_alive():
logger.bind(tag=TAG).warning(
"TTS工作线程未能按时退出: {}", thread.name
)
self._sentence_text_map.clear()
if hasattr(self, "ws") and self.ws:
await self.ws.close()