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:y6aHwySAhzrbibLD@222.186.10.253: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(' len(data_bytes): break value, = struct.unpack(' 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("程序已退出")