Files
banban/talkingq-url/services/interrupt_handler.py
2026-03-24 15:04:36 +08:00

243 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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