from typing import Dict, Optional import time from sqlalchemy import select, update, insert, delete from database.models import DeviceFirmwareUpdate from services.database_service_base import DatabaseServiceBase from utils.logger import session_logger class DeviceFirmwareUpdateManager(DatabaseServiceBase): def __init__(self): super().__init__(service_name="device_firmware_update") self.update_cache = {} # 添加缓存以减少数据库查询 self.cache_expiry = 300 # 缓存过期时间(秒) self.cache_timestamps = {} # 记录缓存更新时间 self.max_cache_size = 1000 # 最大缓存条目数 async def create_firmware_update(self, device_id: str, firmware_version: str, update_status: str = "success"): """新增固件更新记录""" await self._init_database() db_session = await self.get_session() try: query = select(DeviceFirmwareUpdate).where(DeviceFirmwareUpdate.device_id == device_id) result = await db_session.execute(query) existing = result.scalar_one_or_none() if existing: update_stmt = ( update(DeviceFirmwareUpdate) .where(DeviceFirmwareUpdate.device_id == device_id) .values( firmware_version=firmware_version, update_status=update_status, progress=0.0 ) ) await db_session.execute(update_stmt) else: insert_stmt = insert(DeviceFirmwareUpdate).values( device_id=device_id, firmware_version=firmware_version, update_status=update_status, progress=0.0 ) await db_session.execute(insert_stmt) await db_session.commit() self._update_cache(device_id, { "device_id": device_id, "firmware_version": firmware_version, "update_status": update_status, "progress": 0.0 }) return True except Exception as e: await db_session.rollback() session_logger.error(device_id, "firmware_update", f"创建固件更新记录失败: {str(e)}") return False finally: await db_session.close() async def get_firmware_update(self, device_id: str) -> Optional[DeviceFirmwareUpdate]: """获取设备固件更新信息""" cached_data = self._get_from_cache(device_id) if cached_data: return cached_data await self._init_database() db_session = await self.get_session() try: query = select(DeviceFirmwareUpdate).where(DeviceFirmwareUpdate.device_id == device_id) result = await db_session.execute(query) update_record = result.scalars().first() if update_record: self._update_cache(device_id, update_record) return update_record except Exception as e: session_logger.error(device_id, "firmware_update", f"查询固件更新失败: {str(e)}") return None finally: await db_session.close() async def update_firmware_update(self, device_id: str, firmware_version: str, update_status: str) -> bool: """更新设备固件信息(支持部分字段更新)""" await self._init_database() db_session = await self.get_session() try: update_values = { "firmware_version": firmware_version, "update_status": update_status } query = select(DeviceFirmwareUpdate).where(DeviceFirmwareUpdate.device_id == device_id) result = await db_session.execute(query) device_update = result.scalars().first() if device_update: update_stmt = ( update(DeviceFirmwareUpdate) .where(DeviceFirmwareUpdate.device_id == device_id) .values(**update_values) ) await db_session.execute(update_stmt) else: insert_stmt = insert(DeviceFirmwareUpdate).values( device_id=device_id, firmware_version=firmware_version, update_status=update_status ) await db_session.execute(insert_stmt) await db_session.commit() self._clear_cache(device_id) return True except Exception as e: await db_session.rollback() session_logger.error(device_id, "firmware_update", f"更新固件信息失败: {str(e)}") return False finally: await db_session.close() async def update_firmware_progress(self, device_id: str, progress: float) -> bool: """更新升级进度""" await self._init_database() db_session = await self.get_session() try: update_stmt = ( update(DeviceFirmwareUpdate) .where(DeviceFirmwareUpdate.device_id == device_id) .values(progress=progress) ) await db_session.execute(update_stmt) await db_session.commit() if device_id in self.update_cache: self.update_cache[device_id].progress = progress self.cache_timestamps[device_id] = time.time() session_logger.info(device_id, "firmware_update", f"更新进度已更新: {progress:.1f}%") return True except Exception as e: await db_session.rollback() session_logger.error(device_id, "firmware_update", f"更新固件进度失败: {str(e)}") return False finally: await db_session.close() async def delete_firmware_update(self, device_id: str) -> bool: """删除设备固件更新记录""" await self._init_database() db_session = await self.get_session() try: delete_stmt = delete(DeviceFirmwareUpdate).where(DeviceFirmwareUpdate.device_id == device_id) await db_session.execute(delete_stmt) await db_session.commit() if device_id in self.update_cache: del self.update_cache[device_id] if device_id in self.cache_timestamps: del self.cache_timestamps[device_id] session_logger.info(device_id, "firmware_update", "已删除固件更新记录") return True except Exception as e: await db_session.rollback() session_logger.error(device_id, "firmware_update", f"删除固件更新记录失败: {str(e)}") return False finally: await db_session.close() async def update_device_firmware_version(self, device_id: str, firmware_version: str) -> bool: """更新设备当前固件版本信息""" await self._init_database() db_session = await self.get_session() try: query = select(DeviceFirmwareUpdate).where(DeviceFirmwareUpdate.device_id == device_id) result = await db_session.execute(query) device_update = result.scalars().first() if device_update: update_stmt = ( update(DeviceFirmwareUpdate) .where(DeviceFirmwareUpdate.device_id == device_id) .values(firmware_version=firmware_version) ) await db_session.execute(update_stmt) else: insert_stmt = insert(DeviceFirmwareUpdate).values( device_id=device_id, firmware_version=firmware_version, update_status="success" # 默认状态为成功 ) await db_session.execute(insert_stmt) await db_session.commit() self._clear_cache(device_id) return True except Exception as e: await db_session.rollback() session_logger.error(device_id, "firmware_update", f"更新固件版本失败: {str(e)}") return False finally: await db_session.close() async def request_firmware_version(self, device_id: str) -> bool: """向设备请求当前固件版本""" from services.connection_manager import connection_manager websocket = await connection_manager.get_connection(device_id) if not websocket: session_logger.error(device_id, "firmware_update", f"设备 {device_id} 未在线或未找到") return False try: await websocket.send_text("GET_FIRMWARE_VERSION") session_logger.info(device_id, "firmware_update", "已向设备发送固件版本请求") return True except Exception as e: session_logger.error(device_id, "firmware_update", f"请求设备固件版本失败: {str(e)}") return False def _update_cache(self, device_id: str, update_record): """更新缓存""" # 检查缓存大小限制 if len(self.update_cache) >= self.max_cache_size: self._cleanup_old_cache_entries() self.update_cache[device_id] = update_record self.cache_timestamps[device_id] = time.time() def _cleanup_old_cache_entries(self): """清理最旧的缓存条目""" if not self.cache_timestamps: return # 按时间戳排序,删除最旧的25%条目 sorted_items = sorted(self.cache_timestamps.items(), key=lambda x: x[1]) cleanup_count = max(1, len(sorted_items) // 4) for device_id, _ in sorted_items[:cleanup_count]: if device_id in self.update_cache: del self.update_cache[device_id] if device_id in self.cache_timestamps: del self.cache_timestamps[device_id] session_logger.info( "system", "cache_cleanup", f"设备固件缓存清理完成,删除了 {cleanup_count} 个条目,剩余 {len(self.update_cache)} 个" ) def _get_from_cache(self, device_id: str): """从缓存获取记录,如果缓存过期则返回None""" if device_id in self.update_cache and device_id in self.cache_timestamps: if time.time() - self.cache_timestamps[device_id] < self.cache_expiry: return self.update_cache[device_id] else: del self.update_cache[device_id] del self.cache_timestamps[device_id] return None def _clear_cache(self, device_id: str): """清除设备的缓存""" if device_id in self.update_cache: del self.update_cache[device_id] if device_id in self.cache_timestamps: del self.cache_timestamps[device_id] device_firmware_update_manager = DeviceFirmwareUpdateManager()