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