685 lines
27 KiB
Python
685 lines
27 KiB
Python
import socket
|
|
import asyncio
|
|
import json
|
|
from motor.motor_asyncio import AsyncIOMotorClient
|
|
from typing import Dict, List
|
|
import struct
|
|
from bson.objectid import ObjectId
|
|
import time
|
|
from redis import asyncio as aioredis
|
|
import websockets
|
|
import signal
|
|
|
|
|
|
class UDPServer(asyncio.DatagramProtocol):
|
|
def __init__(self, host='0.0.0.0', port=6002):
|
|
self.server_address = (host, port)
|
|
self.loop = asyncio.get_event_loop()
|
|
self.transport = None
|
|
|
|
# MongoDB 连接
|
|
self.mongodb_url = "mongodb://lab:[email protected]:27017/lab"
|
|
self.db_client = None
|
|
self.redis_client = None
|
|
|
|
# 缓存活跃实验的设备映射
|
|
self.active_experiments: Dict[str, dict] = {} # serial_number -> {experiment_id, sensors_map}
|
|
self.active_experiments_lock = asyncio.Lock() # 添加锁来保护活跃实验映射
|
|
|
|
# WebSocket连接
|
|
self.websocket: Dict[str, websockets.WebSocketClientProtocol] = {} # serial_number -> WebSocketClientProtocol
|
|
self.experiment_id: Dict[str, str] = {} # serial_number -> experiment_id
|
|
self.reconnect_lock = asyncio.Lock() # 添加重连锁
|
|
self.last_connect_attempt = 0 # 添加最后连接尝试时间
|
|
self.RECONNECT_COOLDOWN = 30 # 重连冷却时间(秒)
|
|
self.MAX_RECONNECT_ATTEMPTS = 3 # 最大重连次数
|
|
self.reconnect_attempts = 0 # 重连计数器
|
|
|
|
# 添加设备序列号映射
|
|
self.device_serials: Dict[tuple, str] = {} # (ip, port) -> serial_number
|
|
|
|
# 添加重试相关的属性
|
|
self.serial_request_intervals = {} # (ip, port) -> last_request_time
|
|
self.experiment_check_intervals = {} # serial_number -> last_check_time
|
|
self.REQUEST_INTERVAL = 1 # 请求序列号的间隔(秒)
|
|
self.EXPERIMENT_CHECK_INTERVAL = 5 # 每5秒检查一次
|
|
|
|
# 添加时间戳相关属性
|
|
self.last_timestamp = None
|
|
self.timestamp_offset = 0
|
|
|
|
# 添加实验更新相关的属性
|
|
self.last_experiment_update = 0
|
|
self.EXPERIMENT_UPDATE_INTERVAL = 1 # 实验更新间隔(秒)
|
|
|
|
self.SAMPLE_INTERVAL_MS = 1 # 修改为1毫秒,而不是之前的4毫秒
|
|
|
|
self._shutdown = False
|
|
self._shutdown_event = asyncio.Event()
|
|
|
|
# 添加定时查询相关的属性
|
|
self.experiment_check_task = None
|
|
|
|
# 添加Redis批量处理
|
|
self.redis_batch_size = 100
|
|
self.redis_batches: Dict[str, List] = {}
|
|
self.redis_batch_locks: Dict[str, asyncio.Lock] = {}
|
|
|
|
# 添加实验状态缓存
|
|
self.experiment_status: Dict[str, Dict[str, str]] = {} # experiment_id -> {serial_number: status}
|
|
self.status_check_interval = 5 # 每5秒检查一次状态
|
|
|
|
# 添加状态消息缓存
|
|
self.status_message_cache: Dict[str, bool] = {} # experiment_id:serial_number -> has_printed
|
|
|
|
print(f"UDP服务器启动在 {host}:{port}")
|
|
|
|
async def init_mongodb(self):
|
|
"""初始化MongoDB连接"""
|
|
if not self.db_client:
|
|
self.db_client = AsyncIOMotorClient(self.mongodb_url)
|
|
|
|
async def init_redis(self):
|
|
"""初始化Redis连接"""
|
|
if not self.redis_client:
|
|
self.redis_client = await aioredis.from_url(
|
|
'redis://222.186.10.253:6379',
|
|
password='Obscura@2024',
|
|
db=200,
|
|
decode_responses=True
|
|
)
|
|
|
|
async def update_active_experiments(self):
|
|
"""增量更新活跃实验的设备映射"""
|
|
if not self.db_client:
|
|
await self.init_mongodb()
|
|
db = self.db_client["lab"]
|
|
|
|
try:
|
|
async with self.active_experiments_lock:
|
|
# print("开始查询活跃实验...")
|
|
|
|
# 查找所有活跃的实验会话
|
|
cursor = db.experiment_sessions.find({"end_time": None})
|
|
count = await db.experiment_sessions.count_documents({"end_time": None})
|
|
# print(f"找到 {count} 个活跃实验会话")
|
|
|
|
# 保留当前的映射
|
|
current_mapping = self.active_experiments.copy()
|
|
|
|
# 增量更新映射
|
|
async for session in cursor:
|
|
experiment_id = str(session["experiment_id"])
|
|
# print(f"处理实验会话: {experiment_id}")
|
|
|
|
for device in session.get("devices", []):
|
|
serial_number = device.get("serial_number")
|
|
if serial_number:
|
|
# print(f"检查设备: {serial_number}")
|
|
|
|
# 构建传感器映射
|
|
sensors_map = {}
|
|
for sensor in device.get("sensors", []):
|
|
channel_index = sensor.get("index")
|
|
if channel_index and channel_index.startswith("channel"):
|
|
channel_num = int(channel_index[7:]) - 1
|
|
sensors_map[channel_num] = sensor.get("sensor_name")
|
|
|
|
# 更新映射
|
|
current_mapping[serial_number] = {
|
|
"experiment_id": experiment_id,
|
|
"sensors_map": sensors_map
|
|
}
|
|
# print(f"更新设备映射: {serial_number} -> {experiment_id}")
|
|
|
|
# 检查映射是否有变化
|
|
if current_mapping != self.active_experiments:
|
|
# print(f"活跃实验映射已更新: {current_mapping}")
|
|
self.active_experiments = current_mapping
|
|
return True
|
|
else:
|
|
# print("没有发现新的实验映射")
|
|
return False
|
|
|
|
except Exception as e:
|
|
print(f"更新活跃实验时出错: {e}")
|
|
import traceback
|
|
traceback.print_exc()
|
|
return False
|
|
|
|
def connection_made(self, transport):
|
|
self.transport = transport
|
|
print(f"UDP服务器启动在 {self.server_address[0]}:{self.server_address[1]}")
|
|
# 启动定时查询任务
|
|
self.experiment_check_task = asyncio.create_task(self.periodic_experiment_check())
|
|
|
|
async def periodic_experiment_check(self):
|
|
"""定期检查活跃实验"""
|
|
while not self._shutdown:
|
|
try:
|
|
# print("执行定期实验检查...")
|
|
await self.update_active_experiments()
|
|
await asyncio.sleep(self.EXPERIMENT_CHECK_INTERVAL)
|
|
except Exception as e:
|
|
print(f"定期检查实验时出错: {e}")
|
|
await asyncio.sleep(1) # 出错时短暂等待后重试
|
|
|
|
def datagram_received(self, data, addr):
|
|
"""处理接收到的UDP数据"""
|
|
try:
|
|
# 检查是否是序列号
|
|
try:
|
|
text = data.decode('utf-8').strip()
|
|
if text and all(c.isprintable() for c in text):
|
|
print(f"检测到序列号: {text}")
|
|
asyncio.create_task(self.handle_serial_number(text, addr))
|
|
return
|
|
except UnicodeDecodeError:
|
|
pass
|
|
|
|
# 处理传感器数据
|
|
task = asyncio.create_task(self.handle_sensor_data(data, addr))
|
|
# 添加任务完成回调
|
|
task.add_done_callback(self._handle_task_result)
|
|
|
|
except Exception as e:
|
|
print(f"数据处理错误: {e}")
|
|
|
|
def _handle_task_result(self, task):
|
|
"""处理任务完成的回调"""
|
|
try:
|
|
# 获取任务结果,如果有异常会在这里抛出
|
|
task.result()
|
|
except asyncio.CancelledError:
|
|
pass # 任务被取消是正常的
|
|
except Exception as e:
|
|
print(f"任务执行出错: {e}")
|
|
import traceback
|
|
traceback.print_exc()
|
|
|
|
def request_serial_number(self, addr):
|
|
"""向设备请求序列号"""
|
|
try:
|
|
# 修改为与 boot.py 中匹配的请求命令
|
|
self.transport.sendto(b'REQUEST_SERIAL', addr)
|
|
except Exception as e:
|
|
print(f"请求序列号失败: {e}")
|
|
|
|
async def decode_micropython_data(self, data, addr):
|
|
try:
|
|
# 包头部信息(前4字节)
|
|
timestamp, chunk_index = struct.unpack('<HH', data[:4])
|
|
|
|
# 获取设备的序列号
|
|
serial_number = self.device_serials.get(addr)
|
|
if not serial_number:
|
|
print(f"收到未知设备数据,请求序列号: {addr}")
|
|
self.request_serial_number(addr)
|
|
return None
|
|
|
|
# 使用锁保护读取操作
|
|
async with self.active_experiments_lock:
|
|
experiment_info = self.active_experiments.get(serial_number)
|
|
if not experiment_info:
|
|
# print(f"设备 {serial_number} 未找到实验,等待下次检查...")
|
|
return None
|
|
|
|
# 取通道配置
|
|
sensors_map = experiment_info.get("sensors_map", {})
|
|
if not sensors_map:
|
|
print(f"设备 {serial_number} 的传感器映射为空")
|
|
return None
|
|
|
|
num_channels = len(sensors_map)
|
|
if num_channels == 0:
|
|
print(f"设备 {serial_number} 的通道配置为空")
|
|
return None
|
|
|
|
# 解析传感器数据
|
|
sensor_data = []
|
|
data_bytes = data[4:] # 跳过头部4字节
|
|
|
|
# 每2字节解析一个数值,不再使用0值作为停止条件
|
|
for i in range(0, len(data_bytes), 2):
|
|
if i + 2 > len(data_bytes):
|
|
break
|
|
value, = struct.unpack('<H', data_bytes[i:i+2])
|
|
sensor_data.append(value)
|
|
|
|
# 检查数据长度是否是通道数的整数倍
|
|
if len(sensor_data) % num_channels != 0:
|
|
print(f"设备 {serial_number} 的数据长度 {len(sensor_data)} 不是通道数 {num_channels} 的整数倍")
|
|
return None
|
|
|
|
# 将数据重组为对应数的通道
|
|
channels = [[] for _ in range(num_channels)]
|
|
for i in range(0, len(sensor_data), num_channels):
|
|
for ch in range(num_channels):
|
|
channels[ch].append(sensor_data[i + ch])
|
|
|
|
# 检查每个通道是否都有数据
|
|
if not all(len(ch) > 0 for ch in channels):
|
|
print(f"设备 {serial_number} 的某些通道没有数据")
|
|
return None
|
|
|
|
return {
|
|
'timestamp': timestamp,
|
|
'chunk_index': chunk_index,
|
|
'channels': channels,
|
|
'sample_count': len(channels[0]),
|
|
'serial_number': serial_number
|
|
}
|
|
|
|
except Exception as e:
|
|
print(f"数据解析错误: {e}")
|
|
import traceback
|
|
traceback.print_exc()
|
|
return None
|
|
|
|
async def ensure_websocket_connected(self, experiment_id: str, serial_number: str):
|
|
"""确保WebSocket连接保持活跃"""
|
|
try:
|
|
# 检查现有连接是否可用
|
|
if (serial_number in self.websocket and
|
|
self.experiment_id.get(serial_number) == experiment_id and
|
|
not self.websocket[serial_number].closed):
|
|
return True
|
|
|
|
async with self.reconnect_lock:
|
|
# 关闭旧连接
|
|
await self.close_websocket(serial_number)
|
|
|
|
uri = f"ws://dev.obscura.work/lab/ws/{experiment_id}/{serial_number}"
|
|
print(f"建立新的WebSocket连接: {uri}")
|
|
|
|
try:
|
|
websocket = await websockets.connect(
|
|
uri,
|
|
ping_interval=20, # 添加心跳机制
|
|
ping_timeout=10,
|
|
close_timeout=5
|
|
)
|
|
|
|
# 发送初始连接消息
|
|
init_message = {
|
|
"type": "connect",
|
|
"serial_number": serial_number,
|
|
"experiment_id": experiment_id
|
|
}
|
|
await websocket.send(json.dumps(init_message))
|
|
response = await websocket.recv()
|
|
response_data = json.loads(response)
|
|
|
|
if response_data.get("status") == "success":
|
|
self.websocket[serial_number] = websocket
|
|
self.experiment_id[serial_number] = experiment_id
|
|
print(f"WebSocket连接成功建立: {serial_number}")
|
|
return True
|
|
|
|
await websocket.close()
|
|
return False
|
|
|
|
except Exception as e:
|
|
print(f"WebSocket连接失败: {e}")
|
|
return False
|
|
|
|
except Exception as e:
|
|
print(f"确保WebSocket连接时出错: {e}")
|
|
return False
|
|
|
|
async def close_websocket(self, serial_number):
|
|
"""安全关闭特定设备的WebSocket连接"""
|
|
if serial_number in self.websocket:
|
|
try:
|
|
await self.websocket[serial_number].close()
|
|
except Exception:
|
|
pass
|
|
del self.websocket[serial_number]
|
|
del self.experiment_id[serial_number]
|
|
|
|
async def batch_redis_write(self, stream_key: str, data_point: dict):
|
|
"""批量写入Redis"""
|
|
if stream_key not in self.redis_batches:
|
|
self.redis_batches[stream_key] = []
|
|
self.redis_batch_locks[stream_key] = asyncio.Lock()
|
|
|
|
async with self.redis_batch_locks[stream_key]:
|
|
self.redis_batches[stream_key].append(data_point)
|
|
|
|
if len(self.redis_batches[stream_key]) >= self.redis_batch_size:
|
|
batch = self.redis_batches[stream_key]
|
|
self.redis_batches[stream_key] = []
|
|
|
|
# 使用pipeline批量写入
|
|
pipe = self.redis_client.pipeline()
|
|
for point in batch:
|
|
pipe.xadd(stream_key, point)
|
|
await pipe.execute()
|
|
|
|
async def check_experiment_status(self, experiment_id: str, serial_number: str) -> bool:
|
|
"""检查实验状态"""
|
|
try:
|
|
# 确保Redis连接已初始化
|
|
if not self.redis_client:
|
|
await self.init_redis()
|
|
|
|
# 切换到db199
|
|
await self.redis_client.select(199)
|
|
|
|
status_key = f"experiment_status:{experiment_id}:{serial_number}"
|
|
status = await self.redis_client.get(status_key)
|
|
|
|
# 切回db200(用于其他操作)
|
|
await self.redis_client.select(200)
|
|
|
|
return status == "active"
|
|
except Exception as e:
|
|
print(f"检查实验状态时出错: {e}")
|
|
return False
|
|
|
|
async def process_and_forward_data(self, parsed_data):
|
|
try:
|
|
serial_number = parsed_data.get('serial_number')
|
|
|
|
if not serial_number or serial_number not in self.active_experiments:
|
|
print(f"序列号 {serial_number} 未在活跃实验中")
|
|
return
|
|
|
|
experiment_info = self.active_experiments[serial_number]
|
|
experiment_id = experiment_info["experiment_id"]
|
|
sensors_map = experiment_info["sensors_map"]
|
|
|
|
# 确保Redis连接已初始化
|
|
if not self.redis_client:
|
|
await self.init_redis()
|
|
|
|
channels = parsed_data['channels']
|
|
timestamp = parsed_data['timestamp']
|
|
|
|
# Redis stream 处理 - 无论实验状态如何都保存数据
|
|
stream_key = f"experiment:{experiment_id}:{serial_number}"
|
|
channel_data = {
|
|
sensor_name: channels[channel_num]
|
|
for channel_num, sensor_name in sensors_map.items()
|
|
if channel_num < len(channels)
|
|
}
|
|
|
|
# Redis批量写入
|
|
await self.batch_redis_write(stream_key, {
|
|
'data': json.dumps({
|
|
'channel_data': channel_data,
|
|
'timestamp': timestamp,
|
|
'serial_number': serial_number,
|
|
'total_points': parsed_data['sample_count']
|
|
})
|
|
})
|
|
|
|
# 检查实验是否处于活跃状态
|
|
is_active = await self.check_experiment_status(experiment_id, serial_number)
|
|
if not is_active:
|
|
# 使用组合键检查是否已经打印过停止消息
|
|
cache_key = f"{experiment_id}:{serial_number}"
|
|
if not self.status_message_cache.get(cache_key):
|
|
print(f"实验ID: {experiment_id}, 设备: {serial_number} 已停止,仅保存数据")
|
|
self.status_message_cache[cache_key] = True
|
|
return
|
|
else:
|
|
# 如果实验重新激活,清除缓存
|
|
cache_key = f"{experiment_id}:{serial_number}"
|
|
self.status_message_cache.pop(cache_key, None)
|
|
|
|
# 以下是WebSocket相关操作,只在实验活跃时执行
|
|
websocket = self.websocket.get(serial_number)
|
|
if not websocket or websocket.closed:
|
|
print(f"重新建立WebSocket连接: {serial_number}")
|
|
connected = await self.ensure_websocket_connected(experiment_id, serial_number)
|
|
if connected:
|
|
websocket = self.websocket.get(serial_number)
|
|
else:
|
|
print(f"无法建立WebSocket连接: {serial_number}")
|
|
return
|
|
|
|
try:
|
|
# 降采样处理
|
|
data_length = len(next(iter(channel_data.values())))
|
|
DOWNSAMPLE_RATE = max(5, data_length // 1000) # 动态调整降采样率
|
|
|
|
# 使用numpy进行高效降采样
|
|
import numpy as np
|
|
downsampled_data = {}
|
|
for sensor_name, data in channel_data.items():
|
|
data_array = np.array(data)
|
|
downsampled_data[sensor_name] = data_array[::DOWNSAMPLE_RATE].tolist()
|
|
|
|
# 计算时间戳序列
|
|
base_timestamp = timestamp
|
|
sample_interval = self.SAMPLE_INTERVAL_MS * DOWNSAMPLE_RATE
|
|
timestamps = [
|
|
base_timestamp + (i * sample_interval)
|
|
for i in range(len(next(iter(downsampled_data.values()))))
|
|
]
|
|
|
|
# 构建批量数据点
|
|
data_points = []
|
|
for i in range(len(timestamps)):
|
|
point_data = {
|
|
sensor_name: values[i]
|
|
for sensor_name, values in downsampled_data.items()
|
|
if i < len(values)
|
|
}
|
|
data_points.append({
|
|
'sensor_data': point_data,
|
|
'timestamp': timestamps[i]
|
|
})
|
|
|
|
# 批量发送数据
|
|
if websocket and not websocket.closed:
|
|
message = {
|
|
'type': 'data_batch',
|
|
'data_points': data_points,
|
|
'serial_number': serial_number,
|
|
'total_points': len(data_points)
|
|
}
|
|
# print(f"发送数据: {message}")
|
|
await websocket.send(json.dumps(message))
|
|
|
|
except websockets.exceptions.ConnectionClosed:
|
|
print(f"发送数据时连接断开,将在下次发送时重连: {serial_number}")
|
|
await self.close_websocket(serial_number)
|
|
|
|
except Exception as e:
|
|
print(f"数据处理错误: {e}")
|
|
import traceback
|
|
traceback.print_exc()
|
|
|
|
except Exception as e:
|
|
print(f"数据处理错误: {e}")
|
|
import traceback
|
|
traceback.print_exc()
|
|
|
|
async def handle_serial_number(self, serial_number: str, addr):
|
|
"""处理接收的序列号"""
|
|
try:
|
|
# 检查是否是的序列号
|
|
is_new_serial = addr not in self.device_serials or self.device_serials[addr] != serial_number
|
|
|
|
# 更新设备地址映射
|
|
self.device_serials[addr] = serial_number
|
|
self.serial_request_intervals.pop(addr, None)
|
|
|
|
# 只在新序列号且未匹配到实验时才更新
|
|
if is_new_serial and serial_number not in self.active_experiments:
|
|
print(f"序列号 {serial_number} 未匹配到实验,正在更新...")
|
|
await self.update_active_experiments()
|
|
if serial_number in self.active_experiments:
|
|
experiment_id = self.active_experiments[serial_number]["experiment_id"]
|
|
print(f"序列号 {serial_number} 已匹配到实验 {experiment_id}")
|
|
else:
|
|
print(f"序列号 {serial_number} 未找到匹配的活跃实验")
|
|
|
|
except Exception as e:
|
|
print(f"处理序列号错误: {e}")
|
|
|
|
async def handle_sensor_data(self, data, addr):
|
|
"""处理传感器数据"""
|
|
try:
|
|
# 解析数据
|
|
parsed_data = await self.decode_micropython_data(data, addr)
|
|
if not parsed_data:
|
|
return
|
|
|
|
# 处理和转发数据
|
|
await self.process_and_forward_data(parsed_data)
|
|
|
|
except asyncio.CancelledError:
|
|
# 优雅地处理取消
|
|
print("传感器数据处理被取消")
|
|
raise
|
|
except Exception as e:
|
|
print(f"处理传感器数据时出错: {e}")
|
|
import traceback
|
|
traceback.print_exc()
|
|
finally:
|
|
# 确保在任务结束时进行清理
|
|
pass
|
|
|
|
def start(self):
|
|
"""启动UDP服务器"""
|
|
print("正在启动UDP服务器...")
|
|
# 初始化数据库连接
|
|
self.loop.run_until_complete(self.init_mongodb())
|
|
# 初始化Redis接
|
|
self.loop.run_until_complete(self.init_redis())
|
|
# 更新活跃实验
|
|
self.loop.run_until_complete(self.update_active_experiments())
|
|
|
|
listen = self.loop.create_datagram_endpoint(
|
|
lambda: self,
|
|
local_addr=self.server_address
|
|
)
|
|
self.transport, _ = self.loop.run_until_complete(listen)
|
|
try:
|
|
self.loop.run_forever()
|
|
except KeyboardInterrupt:
|
|
print("\n服务器正在关闭...")
|
|
finally:
|
|
if self.redis_client:
|
|
self.redis_client.close()
|
|
if self.transport:
|
|
self.transport.close()
|
|
self.loop.close()
|
|
|
|
def stop(self):
|
|
"""停止服务器"""
|
|
print("\n服务器正在关闭...")
|
|
if self.websocket:
|
|
asyncio.create_task(self.websocket.close())
|
|
if self.transport:
|
|
self.transport.close()
|
|
self.loop.close()
|
|
|
|
async def close(self):
|
|
"""关闭服务器并清理资源"""
|
|
try:
|
|
# 关闭所有 WebSocket 连接
|
|
if self.websocket:
|
|
# 创建所有序列号的副本,因为在关闭过程中字典会被修改
|
|
serial_numbers = list(self.websocket.keys())
|
|
for serial_number in serial_numbers:
|
|
await self.close_websocket(serial_number)
|
|
|
|
# 关闭数据库连接
|
|
if self.db_client:
|
|
self.db_client.close()
|
|
|
|
# 关闭 Redis 连接
|
|
if self.redis_client:
|
|
await self.redis_client.close()
|
|
|
|
# 关闭传输层
|
|
if self.transport:
|
|
self.transport.close()
|
|
|
|
except Exception as e:
|
|
print(f"关闭服务器时出错: {e}")
|
|
|
|
async def shutdown(self):
|
|
"""优雅关闭服务器"""
|
|
if self._shutdown:
|
|
return
|
|
|
|
print("\n正在关闭服务器...")
|
|
self._shutdown = True
|
|
|
|
# 取消定时查询任务
|
|
if self.experiment_check_task:
|
|
self.experiment_check_task.cancel()
|
|
try:
|
|
await self.experiment_check_task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
|
|
# 关闭所有连接和资源
|
|
try:
|
|
# 关闭所有 WebSocket 连接
|
|
if self.websocket:
|
|
serial_numbers = list(self.websocket.keys())
|
|
for serial_number in serial_numbers:
|
|
await self.close_websocket(serial_number)
|
|
|
|
# 关闭数据库连接
|
|
if self.db_client:
|
|
self.db_client.close()
|
|
|
|
# 关闭 Redis 连接
|
|
if self.redis_client:
|
|
await self.redis_client.close()
|
|
|
|
# 关闭传输层
|
|
if self.transport:
|
|
self.transport.close()
|
|
|
|
except Exception as e:
|
|
print(f"关闭服务器时出错: {e}")
|
|
finally:
|
|
# 确保事件循环停止
|
|
loop = asyncio.get_event_loop()
|
|
loop.stop()
|
|
|
|
async def main():
|
|
"""主函数"""
|
|
server = UDPServer()
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def signal_handler():
|
|
"""信号处理函数"""
|
|
print("\n收到终止信号")
|
|
asyncio.create_task(server.shutdown())
|
|
|
|
# 注册信号处理器
|
|
for sig in (signal.SIGINT, signal.SIGTERM):
|
|
loop.add_signal_handler(sig, signal_handler)
|
|
|
|
try:
|
|
# 创建服务器
|
|
transport, _ = await loop.create_datagram_endpoint(
|
|
lambda: server,
|
|
local_addr=server.server_address
|
|
)
|
|
|
|
# 等待关闭信号
|
|
await server._shutdown_event.wait()
|
|
|
|
except Exception as e:
|
|
print(f"服务器运行出错: {e}")
|
|
finally:
|
|
# 确保清理资源
|
|
await server.shutdown()
|
|
|
|
if __name__ == "__main__":
|
|
try:
|
|
asyncio.run(main())
|
|
except KeyboardInterrupt:
|
|
print("\n程序被用户中断")
|
|
finally:
|
|
print("程序已退出") |