add banbanmini backend
This commit is contained in:
242
talkingq-url/services/interrupt_handler.py
Normal file
242
talkingq-url/services/interrupt_handler.py
Normal file
@@ -0,0 +1,242 @@
|
||||
import asyncio
|
||||
from typing import List, Callable, Dict, Coroutine, Tuple
|
||||
from utils.logger import session_logger
|
||||
|
||||
|
||||
class InterruptHandler:
|
||||
def __init__(self):
|
||||
self.cleanup_handlers: Dict[
|
||||
tuple, List[Tuple[int, Callable[..., Coroutine]]]
|
||||
] = {}
|
||||
self.interrupt_states: Dict[tuple, bool] = {}
|
||||
self.DEFAULT_PRIORITY = 100
|
||||
self.PRIORITY = {
|
||||
"NETWORK_CONNECTION": 10, # 网络连接最优先关闭
|
||||
"QUEUE_CLEANUP": 20, # 队列清理的优先级
|
||||
"TASK_CANCELLATION": 30, # 任务取消的优先级
|
||||
"DEFAULT": 100, # 默认优先级
|
||||
}
|
||||
self.session_locks: Dict[tuple, asyncio.Lock] = (
|
||||
{}
|
||||
) # 添加锁字典,用于保护中断状态更新
|
||||
self.interrupt_processing: Dict[tuple, bool] = {}
|
||||
self.pending_interrupts: Dict[tuple, int] = {}
|
||||
|
||||
async def get_session_lock(self, session_key: tuple) -> asyncio.Lock:
|
||||
if session_key not in self.session_locks:
|
||||
self.session_locks[session_key] = asyncio.Lock()
|
||||
return self.session_locks[session_key]
|
||||
|
||||
async def register_session(self, session_key: tuple) -> None:
|
||||
lock = await self.get_session_lock(session_key)
|
||||
async with lock:
|
||||
if session_key not in self.cleanup_handlers:
|
||||
self.cleanup_handlers[session_key] = []
|
||||
self.interrupt_states[session_key] = False
|
||||
self.interrupt_processing[session_key] = False # 初始化处理状态
|
||||
self.pending_interrupts[session_key] = 0 # 初始化挂起的中断请求计数
|
||||
device_id, session_id = session_key
|
||||
session_logger.info(device_id, session_id, "已注册中断处理")
|
||||
|
||||
async def register_cleanup_handler(
|
||||
self,
|
||||
session_key: tuple,
|
||||
handler: Callable[..., Coroutine],
|
||||
priority: int = None,
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
if session_key not in self.cleanup_handlers:
|
||||
await self.register_session(session_key)
|
||||
if priority is None:
|
||||
priority = self.DEFAULT_PRIORITY
|
||||
|
||||
async def wrapped_handler():
|
||||
try:
|
||||
await handler(*args, **kwargs)
|
||||
except Exception as e:
|
||||
device_id, session_id = session_key
|
||||
session_logger.error(
|
||||
device_id, session_id, f"执行中断清理处理器出错: {e}"
|
||||
)
|
||||
|
||||
self.cleanup_handlers[session_key].append((priority, wrapped_handler))
|
||||
|
||||
async def handle_interrupt(self, session_key: tuple) -> None:
|
||||
device_id, session_id = session_key
|
||||
await self.set_interrupt_state(session_key, True) # 使用异步方法设置状态
|
||||
lock = await self.get_session_lock(session_key)
|
||||
async with lock:
|
||||
if self.interrupt_processing.get(session_key, False):
|
||||
self.pending_interrupts[session_key] = (
|
||||
self.pending_interrupts.get(session_key, 0) + 1
|
||||
)
|
||||
session_logger.info(
|
||||
device_id,
|
||||
session_id,
|
||||
f"收到连续中断请求,当前有 {self.pending_interrupts[session_key]} 个中断正在等待",
|
||||
)
|
||||
return
|
||||
self.interrupt_processing[session_key] = True
|
||||
handlers_count = len(self.cleanup_handlers.get(session_key, []))
|
||||
try:
|
||||
session_logger.info(
|
||||
device_id, session_id, f"处理中断, 开始执行 {handlers_count} 个清理任务"
|
||||
)
|
||||
if session_key not in self.cleanup_handlers:
|
||||
session_logger.warning(
|
||||
device_id, session_id, "未找到该会话的中断处理器"
|
||||
)
|
||||
return
|
||||
handlers = self.cleanup_handlers[session_key]
|
||||
if not handlers:
|
||||
session_logger.info(device_id, session_id, "该会话没有注册的中断处理器")
|
||||
return
|
||||
|
||||
sorted_handlers = sorted(handlers, key=lambda x: x[0])
|
||||
priority_groups = {}
|
||||
for priority, handler in sorted_handlers:
|
||||
if priority not in priority_groups:
|
||||
priority_groups[priority] = []
|
||||
priority_groups[priority].append(handler)
|
||||
|
||||
for priority in sorted(priority_groups.keys()):
|
||||
group_handlers = priority_groups[priority]
|
||||
group_size = len(group_handlers)
|
||||
session_logger.info(
|
||||
device_id,
|
||||
session_id,
|
||||
f"执行优先级 {priority} 的 {group_size} 个清理任务"
|
||||
)
|
||||
|
||||
timeout = min(1.0 + 0.5 * group_size, 3.0)
|
||||
from services.task_manager import task_manager
|
||||
tasks = []
|
||||
for handler in group_handlers:
|
||||
task = await task_manager.create_task(
|
||||
handler(),
|
||||
device_id=device_id,
|
||||
session_key=session_key,
|
||||
task_type="interrupt_cleanup"
|
||||
)
|
||||
tasks.append(task)
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(*tasks, return_exceptions=True),
|
||||
timeout=timeout
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
session_logger.warning(
|
||||
device_id,
|
||||
session_id,
|
||||
f"优先级 {priority} 的清理任务执行超时({timeout:.1f}秒)"
|
||||
)
|
||||
|
||||
self.cleanup_handlers[session_key] = []
|
||||
session_logger.info(device_id, session_id, "所有中断清理任务已完成")
|
||||
finally:
|
||||
async with lock:
|
||||
self.interrupt_processing[session_key] = False
|
||||
pending_count = self.pending_interrupts.get(session_key, 0)
|
||||
if pending_count > 0:
|
||||
self.pending_interrupts[session_key] = 0
|
||||
session_logger.info(
|
||||
device_id,
|
||||
session_id,
|
||||
f"有 {pending_count} 个连续中断请求,立即处理",
|
||||
)
|
||||
await task_manager.create_task(
|
||||
self.handle_pending_interrupts(session_key, pending_count),
|
||||
device_id=device_id,
|
||||
session_key=session_key,
|
||||
task_type="pending_interrupt"
|
||||
)
|
||||
|
||||
async def handle_pending_interrupts(self, session_key: tuple, count: int) -> None:
|
||||
device_id, session_id = session_key
|
||||
try:
|
||||
session_logger.info(device_id, session_id, f"处理 {count} 个挂起的中断请求")
|
||||
await self.handle_interrupt(session_key)
|
||||
except Exception as e:
|
||||
session_logger.error(device_id, session_id, f"处理挂起中断请求时出错: {e}")
|
||||
|
||||
async def remove_session(self, session_key: tuple) -> None:
|
||||
lock = await self.get_session_lock(session_key)
|
||||
async with lock:
|
||||
if session_key in self.cleanup_handlers:
|
||||
del self.cleanup_handlers[session_key]
|
||||
if session_key in self.interrupt_states:
|
||||
del self.interrupt_states[session_key]
|
||||
if session_key in self.interrupt_processing:
|
||||
del self.interrupt_processing[session_key]
|
||||
if session_key in self.pending_interrupts:
|
||||
del self.pending_interrupts[session_key]
|
||||
if session_key in self.session_locks:
|
||||
del self.session_locks[session_key]
|
||||
device_id, session_id = session_key
|
||||
session_logger.info(device_id, session_id, "已移除会话的中断处理器和状态")
|
||||
|
||||
def is_interrupted(self, session_key: tuple) -> bool:
|
||||
return self.interrupt_states.get(session_key, False)
|
||||
|
||||
def is_processing_interrupt(self, session_key: tuple) -> bool:
|
||||
return self.interrupt_processing.get(session_key, False)
|
||||
|
||||
async def set_interrupt_state(
|
||||
self, session_key: tuple, state: bool
|
||||
) -> None:
|
||||
if session_key not in self.session_locks:
|
||||
await self.register_session(session_key)
|
||||
lock = await self.get_session_lock(session_key)
|
||||
async with lock:
|
||||
old_state = self.interrupt_states.get(session_key, False)
|
||||
if old_state != state:
|
||||
self.interrupt_states[session_key] = state
|
||||
device_id, session_id = session_key
|
||||
status = "中断" if state else "正常"
|
||||
session_logger.info(device_id, session_id, f"会话状态已设置为{status}")
|
||||
else:
|
||||
self.interrupt_states[session_key] = state
|
||||
|
||||
async def notify_client_interrupt_processed(
|
||||
self, session_key: tuple, websocket
|
||||
) -> bool:
|
||||
if not websocket:
|
||||
return False
|
||||
device_id, session_id = session_key
|
||||
max_retries = 3
|
||||
retry_count = 0
|
||||
while retry_count < max_retries:
|
||||
try:
|
||||
if websocket.client_state.name == "CONNECTED" and not getattr(
|
||||
websocket, "_closed", False
|
||||
):
|
||||
await websocket.send_text("INTERRUPT_PROCESSED")
|
||||
session_logger.info(
|
||||
device_id, session_id, "已通知客户端中断处理完成"
|
||||
)
|
||||
return True
|
||||
else:
|
||||
session_logger.warning(
|
||||
device_id, session_id, "WebSocket连接已关闭,无法通知客户端"
|
||||
)
|
||||
return False
|
||||
except Exception as e:
|
||||
retry_count += 1
|
||||
if retry_count < max_retries:
|
||||
session_logger.warning(
|
||||
device_id,
|
||||
session_id,
|
||||
f"通知客户端失败,尝试重试 ({retry_count}/{max_retries}): {e}",
|
||||
)
|
||||
await asyncio.sleep(0.2)
|
||||
else:
|
||||
session_logger.error(
|
||||
device_id, session_id, f"通知客户端失败,已达最大重试次数: {e}"
|
||||
)
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
interrupt_handler = InterruptHandler()
|
||||
Reference in New Issue
Block a user