Files
2025-01-22 07:34:41 +00:00

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("程序已退出")