243 lines
10 KiB
Python
243 lines
10 KiB
Python
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()
|