Files
banban/talkingq-url/services/device_update_manager.py

370 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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()