370 lines
15 KiB
Python
370 lines
15 KiB
Python
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()
|