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()