from typing import Dict, Optional import time from collections.abc import Mapping from sqlalchemy import select, update, insert, delete, text 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 _get_serial_number(self, db_session, device_id: str) -> str: result = await db_session.execute( text("SELECT serial_number FROM device_auth WHERE device_id = :device_id LIMIT 1"), {"device_id": device_id}, ) row = result.mappings().first() if row and row.get("serial_number"): return str(row["serial_number"]) return "" def _row_to_dict(self, row) -> dict: if row is None: return {} if isinstance(row, Mapping): return dict(row) return { "device_id": row.device_id, "firmware_version": row.firmware_version, "ota_channel": getattr(row, "ota_channel", None), "target_version": getattr(row, "target_version", None), "firmware_url": getattr(row, "firmware_url", None), "firmware_id": getattr(row, "firmware_id", None), "source": getattr(row, "source", None), "update_status": row.update_status, "progress": row.progress, "created_at": row.created_at, "updated_at": row.updated_at, } async def create_firmware_update( self, device_id: str, firmware_version: str, update_status: str = "success", progress: float = 0.0, *, ota_channel: str | None = None, target_version: str | None = None, firmware_url: str | None = None, firmware_id: int | None = None, source: str | None = None, ): """新增固件更新记录""" 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=progress, ota_channel=ota_channel, target_version=target_version, firmware_url=firmware_url, firmware_id=firmware_id, source=source, ) ) await db_session.execute(update_stmt) else: serial_number = await self._get_serial_number(db_session, device_id) insert_stmt = insert(DeviceFirmwareUpdate).values( device_id=device_id, serial_number=serial_number, firmware_version=firmware_version, update_status=update_status, progress=progress, ota_channel=ota_channel, target_version=target_version, firmware_url=firmware_url, firmware_id=firmware_id, source=source, ) await db_session.execute(insert_stmt) await db_session.commit() self._update_cache(device_id, { "device_id": device_id, "firmware_version": firmware_version, "ota_channel": ota_channel, "target_version": target_version, "firmware_url": firmware_url, "firmware_id": firmware_id, "source": source, "update_status": update_status, "progress": progress }) 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, progress: Optional[float] = None, *, ota_channel: str | None = None, target_version: str | None = None, firmware_url: str | None = None, firmware_id: int | None = None, source: str | None = None, ) -> bool: """更新设备固件信息(支持部分字段更新)""" await self._init_database() db_session = await self.get_session() try: update_values = { "firmware_version": firmware_version, "update_status": update_status } if progress is not None: update_values["progress"] = progress optional_values = { "ota_channel": ota_channel, "target_version": target_version, "firmware_url": firmware_url, "firmware_id": firmware_id, "source": source, } update_values.update({key: value for key, value in optional_values.items() if value is not None}) 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: serial_number = await self._get_serial_number(db_session, device_id) insert_stmt = insert(DeviceFirmwareUpdate).values( device_id=device_id, serial_number=serial_number, firmware_version=firmware_version, update_status=update_status, progress=progress if progress is not None else 0.0, ota_channel=ota_channel, target_version=target_version, firmware_url=firmware_url, firmware_id=firmware_id, source=source, ) 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() cached_update = self.update_cache.get(device_id) if isinstance(cached_update, dict): cached_update["progress"] = progress elif cached_update is not None: cached_update.progress = progress if cached_update is not None: 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: serial_number = await self._get_serial_number(db_session, device_id) insert_stmt = insert(DeviceFirmwareUpdate).values( device_id=device_id, serial_number=serial_number, 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] async def get_firmware_update_dict(self, device_id: str) -> Optional[dict]: record = await self.get_firmware_update(device_id) if record is None: return None return self._row_to_dict(record) device_firmware_update_manager = DeviceFirmwareUpdateManager()