From be4d2765574896d1d6748f706fc7e7bcf0813997 Mon Sep 17 00:00:00 2001 From: HycJack <772403255@qq.com> Date: Sat, 9 May 2026 09:46:06 +0800 Subject: [PATCH] =?UTF-8?q?=E5=90=88=E5=85=A5=E5=B0=8F=E7=A8=8B=E5=BA=8F?= =?UTF-8?q?=E4=BF=AE=E6=94=B9=EF=BC=8C=E6=96=B0=E5=A2=9E=E5=91=8A=E8=AD=A6?= =?UTF-8?q?=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- talkingq-url/banban/dao/im.py | 46 +++++ talkingq-url/banban/routers/devices.py | 63 ++++++ talkingq-url/banban/service/device.py | 97 +++++++++ talkingq-url/banban/service/im.py | 72 +++++++ talkingq-url/database/models.py | 20 ++ talkingq-url/handlers/mqtt_handler.py | 133 ++++++++++-- talkingq-url/scripts/generate_bind_qr.py | 8 +- .../scripts/import_device_imei_mapping.py | 189 ++++++++++++++++++ .../services/device_identity_initializer.py | 129 ++++++++++++ .../services/device_update_manager.py | 71 ++++++- 10 files changed, 800 insertions(+), 28 deletions(-) create mode 100644 talkingq-url/scripts/import_device_imei_mapping.py create mode 100644 talkingq-url/services/device_identity_initializer.py diff --git a/talkingq-url/banban/dao/im.py b/talkingq-url/banban/dao/im.py index 137c737..4b792a7 100644 --- a/talkingq-url/banban/dao/im.py +++ b/talkingq-url/banban/dao/im.py @@ -18,6 +18,14 @@ class DeviceIdentity: child_name: str | None +@dataclass(frozen=True) +class DeviceOwnerIdentity: + device_id: str + child_id: int + child_name: str | None + owner_user_id: int + + @dataclass(frozen=True) class ConversationMessageCreateResult: idempotent: bool @@ -98,6 +106,44 @@ class ImDAO(BaseDAO): child_id=int(row["child_id"]), child_name=row["child_name"], ) + + async def get_bound_device_owner_identity(self, *, device_id: str) -> DeviceOwnerIdentity: + row = ( + await self.execute( + text( + """ + SELECT + da.device_id, + db.owner_user_id, + db.child_id, + c.child_name + FROM device_auth AS da + JOIN device_bindings AS db + ON db.device_id = da.device_id + AND db.status = 1 + LEFT JOIN children AS c + ON c.child_id = db.child_id + AND c.status = 1 + WHERE da.device_id = :device_id + AND da.is_active = 1 + LIMIT 1 + """ + ), + {"device_id": device_id}, + ) + ).mappings().first() + if not row: + from fastapi import HTTPException, status + raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid device credentials") + if row["child_id"] is None: + from fastapi import HTTPException + raise HTTPException(status_code=404, detail="device not bound to a child") + return DeviceOwnerIdentity( + device_id=str(row["device_id"]), + child_id=int(row["child_id"]), + child_name=row["child_name"], + owner_user_id=int(row["owner_user_id"]), + ) async def get_device_by_id(self, *, device_id: str) -> DeviceIdentity: row = ( diff --git a/talkingq-url/banban/routers/devices.py b/talkingq-url/banban/routers/devices.py index 95b5406..d8336db 100644 --- a/talkingq-url/banban/routers/devices.py +++ b/talkingq-url/banban/routers/devices.py @@ -127,6 +127,21 @@ class DeviceRemoteSleepWakeResponse(BaseModel): msg_id: str +class DeviceFirmwareStatusResponse(BaseModel): + device_id: str + current_version: str | None = None + latest_version: str | None = None + update_available: bool + can_update: bool + update_status: str + progress: float + target_version: str | None = None + updated_at: datetime | None = None + + +class DeviceFirmwareUpdateResponse(DeviceFirmwareStatusResponse): + msg_id: str + @@ -430,6 +445,54 @@ async def set_device_remote_sleep_wake( ) +@router.get("/{device_id}/firmware", response_model=DeviceFirmwareStatusResponse) +async def get_device_firmware_status( + device_id: str, + request: Request, + current_user_id: int = Depends(get_current_user_id), +) -> DeviceFirmwareStatusResponse: + result = await device_service.get_firmware_status( + device_id=device_id, + user_id=current_user_id, + ) + logger.info( + "device firmware status fetched", + extra={ + "event": "device_firmware_status", + "request_id": getattr(request.state, "request_id", None), + "user_id": current_user_id, + "device_id": device_id, + "update_available": result["update_available"], + "update_status": result["update_status"], + }, + ) + return DeviceFirmwareStatusResponse(**result) + + +@router.post("/{device_id}/firmware/update", response_model=DeviceFirmwareUpdateResponse) +async def start_device_firmware_update( + device_id: str, + request: Request, + current_user_id: int = Depends(get_current_user_id), +) -> DeviceFirmwareUpdateResponse: + result = await device_service.start_firmware_update( + device_id=device_id, + user_id=current_user_id, + ) + logger.info( + "device firmware update command sent", + extra={ + "event": "device_firmware_update", + "request_id": getattr(request.state, "request_id", None), + "user_id": current_user_id, + "device_id": device_id, + "target_version": result["target_version"], + "msg_id": result["msg_id"], + }, + ) + return DeviceFirmwareUpdateResponse(**result) + + @router.get("/{device_id}/location", response_model=DeviceLocationCurrentResponse) async def get_current_device_location( diff --git a/talkingq-url/banban/service/device.py b/talkingq-url/banban/service/device.py index 135ca0c..6c28ee3 100644 --- a/talkingq-url/banban/service/device.py +++ b/talkingq-url/banban/service/device.py @@ -7,6 +7,8 @@ from fastapi import HTTPException from banban.dao.device import DeviceDAO from banban.service.device_setting import device_setting_service +from services.device_update_manager import device_firmware_update_manager +from services.system_config_manager import system_config_manager class DeviceService(DatabaseServiceBase): @@ -115,6 +117,101 @@ class DeviceService(DatabaseServiceBase): raise HTTPException(status_code=503, detail="MQTT 服务未初始化") return await service.send_remote_sleep_wake_command(device_id, switch) + def _compare_versions(self, current_version: str | None, latest_version: str | None) -> bool: + current = (current_version or "").strip() + latest = (latest_version or "").strip() + if not latest: + return False + if not current or current in {"unknown", "0.0.0"}: + return True + + try: + current_parts = [int(part) for part in current.split(".")] + latest_parts = [int(part) for part in latest.split(".")] + except ValueError: + return current != latest + + max_len = max(len(current_parts), len(latest_parts)) + current_parts.extend([0] * (max_len - len(current_parts))) + latest_parts.extend([0] * (max_len - len(latest_parts))) + return latest_parts > current_parts + + async def get_firmware_status( + self, + *, + device_id: str, + user_id: int, + ) -> dict[str, Any]: + status_row = await self.get_device_status(device_id=device_id, user_id=user_id) + current_version = status_row.get("version") + + latest_version_config = await system_config_manager.get_config("latest_firmware_version") + firmware_url_config = await system_config_manager.get_config("update_firmware_url") + latest_version = latest_version_config.config_value if latest_version_config else None + firmware_url = firmware_url_config.config_value if firmware_url_config else None + update_available = self._compare_versions(current_version, latest_version) + + update_record = await device_firmware_update_manager.get_firmware_update_dict(device_id) + update_status = update_record.get("update_status") if update_record else "idle" + progress = update_record.get("progress") if update_record else 0.0 + target_version = update_record.get("firmware_version") if update_record else None + updated_at = update_record.get("updated_at") if update_record else None + + return { + "device_id": device_id, + "current_version": current_version, + "latest_version": latest_version, + "update_available": update_available, + "update_status": update_status or "idle", + "progress": float(progress or 0.0), + "target_version": target_version, + "updated_at": updated_at, + "can_update": bool(update_available and latest_version and firmware_url), + } + + async def start_firmware_update( + self, + *, + device_id: str, + user_id: int, + ) -> dict[str, Any]: + firmware_status = await self.get_firmware_status(device_id=device_id, user_id=user_id) + if not firmware_status.get("update_available"): + raise HTTPException(status_code=409, detail="device firmware is already up to date") + + latest_version = firmware_status.get("latest_version") + if not latest_version: + raise HTTPException(status_code=404, detail="latest firmware version not configured") + + firmware_url_config = await system_config_manager.get_config("update_firmware_url") + firmware_url = firmware_url_config.config_value if firmware_url_config else None + if not firmware_url: + raise HTTPException(status_code=404, detail="firmware url not configured") + + from handlers.mqtt_handler import TalkingQMQTTService + + service = await TalkingQMQTTService.get_instance() + if service is None: + raise HTTPException(status_code=503, detail="MQTT 服务未初始化") + + await device_firmware_update_manager.update_firmware_update( + device_id=device_id, + firmware_version=latest_version, + update_status="sent", + progress=0.0, + ) + msg_id = await service.send_ota_command(device_id, firmware_url, latest_version) + + return { + **firmware_status, + "latest_version": latest_version, + "target_version": latest_version, + "update_status": "sent", + "progress": 0.0, + "msg_id": msg_id, + "can_update": bool(firmware_status.get("can_update")), + } + # 创建全局 DeviceService 实例 device_service = DeviceService() diff --git a/talkingq-url/banban/service/im.py b/talkingq-url/banban/service/im.py index f8b732d..9b11433 100644 --- a/talkingq-url/banban/service/im.py +++ b/talkingq-url/banban/service/im.py @@ -71,6 +71,12 @@ def build_device_audio_client_msg_id(*, device_id: str, target_device_id: str, a return f"device-audio-{digest[:32]}" +def build_device_parent_leave_message_client_msg_id(*, device_id: str, media_file_key: str) -> str: + raw = f"{device_id}|parent|{media_file_key}" + digest = hashlib.sha256(raw.encode("utf-8")).hexdigest() + return f"device-parent-{digest[:32]}" + + def normalize_content_json(value: Any) -> dict[str, Any] | None: if value is None: return None @@ -409,6 +415,72 @@ class ImService(DatabaseServiceBase): finally: await db_session.close() + async def create_device_parent_leave_message( + self, + *, + device_id: str, + media_file_key: str, + media_duration_ms: int | None = None, + media_mime_type: str | None = None, + media_size_bytes: int | None = None, + media_transcript_text: str | None = None, + client_msg_id: str | None = None, + ext_json: dict[str, Any] | None = None, + ) -> tuple[DeviceIdentity, ConversationMessageCreateResult]: + normalized_media_file_key = str(media_file_key or "").strip() + if not normalized_media_file_key: + raise HTTPException(status_code=400, detail="media_file_key is required") + + db_session = await self.get_session() + try: + dao = ImDAO(db_session) + owner_identity = await dao.get_bound_device_owner_identity(device_id=device_id) + device_identity = DeviceIdentity( + device_id=owner_identity.device_id, + child_id=owner_identity.child_id, + child_name=owner_identity.child_name, + ) + if ext_json: + resolved_ext_json = dict(ext_json) + else: + resolved_ext_json = {} + resolved_ext_json.update( + { + "message_kind": "leave_message", + "source": "device_mqtt_011", + "source_device_id": device_id, + } + ) + + payload = DeviceMessageCreateRequest( + conversation_type=PARENT_CHILD_CONVERSATION_TYPE, + parent_user_id=owner_identity.owner_user_id, + content_type=2, + media_file_key=normalized_media_file_key, + media_duration_ms=media_duration_ms, + media_mime_type=(media_mime_type or "").strip() or "audio/mpeg", + media_size_bytes=media_size_bytes, + media_transcript_text=(media_transcript_text or "").strip() or None, + client_msg_id=(client_msg_id or "").strip() + or build_device_parent_leave_message_client_msg_id( + device_id=device_id, + media_file_key=normalized_media_file_key, + ), + ext_json=resolved_ext_json, + ) + + result = await self._create_device_message_with_payload( + dao=dao, + device_identity=device_identity, + payload=payload, + ) + return device_identity, result + except Exception: + await db_session.rollback() + raise + finally: + await db_session.close() + async def assert_child_exists(self, *, child_id: int) -> Mapping[str, Any]: db_session = await self.get_session() try: diff --git a/talkingq-url/database/models.py b/talkingq-url/database/models.py index 8b141e6..c9eb83b 100644 --- a/talkingq-url/database/models.py +++ b/talkingq-url/database/models.py @@ -103,10 +103,30 @@ class DeviceAuth(Base): {'mysql_charset': 'utf8mb4', 'mysql_collate': 'utf8mb4_unicode_ci'} ) + +class DeviceImeiMapping(Base): + __tablename__ = "device_imei_mapping" + id = Column(Integer, primary_key=True, autoincrement=True) + imei = Column(String(64), nullable=False, unique=True, index=True) + device_id = Column(String(64), nullable=False, unique=True, index=True) + serial_number = Column(String(64), nullable=False) + status = Column(String(32), nullable=False, server_default=text("'pending'")) + activated_at = Column(DateTime, nullable=True) + created_at = Column(DateTime, nullable=False, server_default=func.now()) + updated_at = Column(DateTime, nullable=False, server_default=func.now(), onupdate=func.now()) + + __table_args__ = ( + Index("idx_device_imei_mapping_status", "status"), + {'mysql_charset': 'utf8mb4', 'mysql_collate': 'utf8mb4_unicode_ci'} + ) + + class DeviceFirmwareUpdate(Base): __tablename__ = "device_firmware_update" id = Column(Integer, primary_key=True, autoincrement=True) device_id = Column(String(64), nullable=False, unique=True, index=True) + serial_number = Column(String(64), nullable=False, default="") + mac_address = Column(String(512), nullable=True) firmware_version = Column(String(64), nullable=False) update_status = Column(String(32), nullable=False, default="success") # 更新状态,如 updating/success/failed created_at = Column(DateTime, nullable=False, server_default=func.now()) diff --git a/talkingq-url/handlers/mqtt_handler.py b/talkingq-url/handlers/mqtt_handler.py index a5801e1..a45481c 100644 --- a/talkingq-url/handlers/mqtt_handler.py +++ b/talkingq-url/handlers/mqtt_handler.py @@ -9,19 +9,19 @@ import aiomqtt from banban.service.binding import BindingService from banban.service.device_alarm import device_alarm_service from banban.service.device_setting import device_setting_service +from banban.service.im import im_service from banban.service.location import location_service from config import settings from services.card_service import card_service -from services.offline_audio_cache import offline_audio_cache -import aiomqtt -from banban.service.location import location_service -from database.models import ChildLocationCurrent -from datetime import datetime from services.device_target_cache import device_target_cache +from services.device_identity_initializer import ( + DeviceIdentityInitializationError, + device_identity_initializer, +) +from services.device_update_manager import device_firmware_update_manager from services.offline_audio_cache import offline_audio_cache from services.task_manager import task_manager from utils.logger import session_logger as logger -from services.task_manager import task_manager class TalkingQMQTTService: _instance = None @@ -55,6 +55,7 @@ class TalkingQMQTTService: "009": self._handle_remote_sleep_wake_response, "010": self._handle_alarm_report, "011": self._handle_short_press_message, + "012": self._handle_device_identity_init, } @classmethod @@ -80,12 +81,16 @@ class TalkingQMQTTService: if not parts or parts[0] != "device": continue - device_id = parts[1] if len(parts) >= 2 else "unknown" - if not device_id.startswith(f"{self.device_prefix}_"): - continue - payload = json.loads(message.payload.decode("utf-8")) msg_id = payload.get("msg_id") + + device_id = parts[1] if len(parts) >= 2 else "unknown" + topic_kind = parts[2] if len(parts) >= 3 else "" + if msg_id == "012" and topic_kind != "event": + continue + if msg_id != "012" and not device_id.startswith(f"{self.device_prefix}_"): + continue + handler = self._msg_handlers.get(msg_id) if handler: await handler(device_id, payload) @@ -185,7 +190,36 @@ class TalkingQMQTTService: async def _handle_ota_response(self, device_id: str, payload: dict): status = payload.get("status") data = payload.get("data", {}) - if status == "success": + progress = data.get("progress") + try: + progress_value = float(progress) if progress is not None else None + except (TypeError, ValueError): + progress_value = None + + if status == "accepted": + await self._schedule_persistence( + device_id, + "ota", + device_firmware_update_manager.update_firmware_update( + device_id=device_id, + firmware_version=data.get("version") or data.get("new_version") or "unknown", + update_status="accepted", + progress=progress_value if progress_value is not None else 0.0, + ), + ) + elif status == "updating": + await self._schedule_persistence( + device_id, + "ota", + device_firmware_update_manager.update_firmware_update( + device_id=device_id, + firmware_version=data.get("version") or data.get("new_version") or "updating", + update_status="updating", + progress=progress_value, + ), + ) + elif status == "success": + new_version = data.get("new_version") or data.get("version") await self._schedule_persistence( device_id, "ota", @@ -194,13 +228,33 @@ class TalkingQMQTTService: power=None, volume=None, signal_strength=None, - version_str=data.get("new_version"), + version_str=new_version, + ), + ) + await self._schedule_persistence( + device_id, + "ota", + device_firmware_update_manager.update_firmware_update( + device_id=device_id, + firmware_version=new_version or "unknown", + update_status="success", + progress=100.0, + ), + ) + elif status == "failed": + await self._schedule_persistence( + device_id, + "ota", + device_firmware_update_manager.update_firmware_update( + device_id=device_id, + firmware_version=data.get("version") or data.get("new_version") or "unknown", + update_status="failed", + progress=progress_value, ), ) - elif status != "accepted": logger.warning(device_id, "", f"[OTA] command failed: {payload}") else: - logger.warning(device_id, "", f"[OTA] 设备 {device_id} 升级异常: {payload}") + logger.warning(device_id, "", f"[OTA] unknown status: {payload}") async def _handle_nfc_notice_response(self, device_id: str, payload: dict): status = payload.get("status") @@ -283,8 +337,35 @@ class TalkingQMQTTService: params = payload.get("params", {}) nfc_uuid = params.get("uuid") logger.info(device_id, "", f"[短按留言] 设备 {device_id} 短按发送留言, UUID={nfc_uuid}") + media_file_key = str(params.get("media_file_key") or params.get("audio_url") or "").strip() + if media_file_key: + try: + await im_service.create_device_parent_leave_message( + device_id=device_id, + media_file_key=media_file_key, + media_duration_ms=params.get("media_duration_ms"), + media_mime_type=params.get("media_mime_type"), + media_size_bytes=params.get("media_size_bytes"), + media_transcript_text=params.get("media_transcript_text"), + client_msg_id=params.get("client_msg_id"), + ext_json=params.get("ext_json") if isinstance(params.get("ext_json"), dict) else None, + ) + logger.info(device_id, "", f"[短按留言] 设备 {device_id} 留言已写入家长会话") + except Exception as exc: + logger.warning(device_id, "", f"[短按留言] 设备 {device_id} 留言写入失败: {exc}") + await self._publish( + f"device/{device_id}/event_resp", + { + "msg_id": "011", + "status": "failed", + "message": str(exc), + }, + ) + return + payload = { "msg_id": "011", + "status": "success", "type": 0, "params": { "url": f"http://{settings.server_host}:{settings.server_port}/assets/audio/message_ok.mp3" @@ -292,6 +373,30 @@ class TalkingQMQTTService: } await self._publish(f"device/{device_id}/event_resp", payload) + async def _handle_device_identity_init(self, imei: str, payload: dict): + try: + result = await device_identity_initializer.initialize_by_imei(imei) + except DeviceIdentityInitializationError as exc: + logger.warning(imei, "", f"[设备初始化] IMEI初始化失败: {exc}") + await self._publish( + f"device/{imei}/event_resp", + { + "msg_id": "012", + "status": "failed", + "message": str(exc), + }, + ) + return + + response_payload = { + "msg_id": "012", + "params": { + "device_id": result.device_id, + "device_sn": result.serial_number, + }, + } + await self._publish(f"device/{imei}/event_resp", response_payload) + async def _send_nfc_listen_response(self, device_id: str, nfc_uuid: str): topic = f"device/{device_id}/event_resp" card = await card_service.get_card_by_uuid(nfc_uuid) diff --git a/talkingq-url/scripts/generate_bind_qr.py b/talkingq-url/scripts/generate_bind_qr.py index d56f4d1..0c63869 100644 --- a/talkingq-url/scripts/generate_bind_qr.py +++ b/talkingq-url/scripts/generate_bind_qr.py @@ -21,7 +21,7 @@ from qrcode.main import QRCode DEFAULT_OUTPUT_DIR = SCRIPT_DIR.parent / "output" / "qr" -SERIAL_PREFIX = "TalkingQ-" +SERIAL_PREFIXES = ("TalkingQ-", "TQ_") SAFE_NAME_RE = re.compile(r"[^A-Za-z0-9._-]+") @@ -32,7 +32,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument("device_id", help="Device ID written into the QR payload.") parser.add_argument( "serial_number", - help=f"Device serial number. It must start with {SERIAL_PREFIX!r}.", + help=f"Device serial number. It must start with one of {SERIAL_PREFIXES!r}.", ) parser.add_argument( "-o", @@ -56,8 +56,8 @@ def validate_inputs(device_id: str, serial_number: str, box_size: int) -> tuple[ raise ValueError("device_id cannot be empty") if not normalized_serial_number: raise ValueError("serial_number cannot be empty") - if not normalized_serial_number.startswith(SERIAL_PREFIX): - raise ValueError(f"serial_number must start with {SERIAL_PREFIX}") + if not normalized_serial_number.startswith(SERIAL_PREFIXES): + raise ValueError(f"serial_number must start with one of {SERIAL_PREFIXES}") if box_size <= 0: raise ValueError("box_size must be greater than 0") diff --git a/talkingq-url/scripts/import_device_imei_mapping.py b/talkingq-url/scripts/import_device_imei_mapping.py new file mode 100644 index 0000000..f6d9d1a --- /dev/null +++ b/talkingq-url/scripts/import_device_imei_mapping.py @@ -0,0 +1,189 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import argparse +import csv +import sys +from pathlib import Path + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +import pymysql +from pymysql.cursors import DictCursor + +from config import settings + + +SERIAL_NUMBER_COLUMNS = ("serial_number", "device_sn") + +CREATE_TABLE_SQL = """ +CREATE TABLE IF NOT EXISTS device_imei_mapping ( + id INT NOT NULL AUTO_INCREMENT, + imei VARCHAR(64) NOT NULL, + device_id VARCHAR(64) NOT NULL, + serial_number VARCHAR(64) NOT NULL, + status VARCHAR(32) NOT NULL DEFAULT 'pending', + activated_at DATETIME NULL, + created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP, + PRIMARY KEY (id), + UNIQUE INDEX imei_UNIQUE (imei ASC), + UNIQUE INDEX device_id_UNIQUE (device_id ASC), + INDEX idx_device_imei_mapping_status (status ASC) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci +""" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Import IMEI to device identity mappings from a CSV file." + ) + parser.add_argument("csv_path", help="CSV file with imei, device_id, serial_number columns.") + parser.add_argument( + "--dry-run", + action="store_true", + help="Validate the CSV without writing to the database.", + ) + return parser.parse_args() + + +def normalize_row(row: dict[str, str], row_number: int) -> dict[str, str]: + serial_number = "" + for column in SERIAL_NUMBER_COLUMNS: + value = str(row.get(column) or "").strip() + if value: + serial_number = value + break + + normalized = { + "imei": str(row.get("imei") or "").strip(), + "device_id": str(row.get("device_id") or "").strip(), + "serial_number": serial_number, + } + missing = [key for key, value in normalized.items() if not value] + if missing: + raise ValueError(f"row {row_number}: missing required columns: {', '.join(missing)}") + return normalized + + +def read_rows(csv_path: Path) -> list[dict[str, str]]: + with csv_path.open("r", encoding="utf-8-sig", newline="") as fp: + reader = csv.DictReader(fp) + fieldnames = set(reader.fieldnames or []) + missing = [column for column in ("imei", "device_id") if column not in fieldnames] + if not any(column in fieldnames for column in SERIAL_NUMBER_COLUMNS): + missing.append("serial_number/device_sn") + if missing: + raise ValueError(f"CSV missing required columns: {', '.join(missing)}") + + rows: list[dict[str, str]] = [] + seen_imei: set[str] = set() + seen_device_id: set[str] = set() + for row_number, row in enumerate(reader, start=2): + normalized = normalize_row(row, row_number) + if normalized["imei"] in seen_imei: + raise ValueError(f"row {row_number}: duplicate imei in CSV: {normalized['imei']}") + if normalized["device_id"] in seen_device_id: + raise ValueError(f"row {row_number}: duplicate device_id in CSV: {normalized['device_id']}") + seen_imei.add(normalized["imei"]) + seen_device_id.add(normalized["device_id"]) + rows.append(normalized) + return rows + + +def get_connection(): + return pymysql.connect( + host=settings.db_host, + port=settings.db_port, + user=settings.db_user, + password=settings.db_password, + database=settings.db_name, + charset="utf8mb4", + cursorclass=DictCursor, + ) + + +def import_rows(rows: list[dict[str, str]]) -> int: + conn = get_connection() + try: + with conn.cursor() as cur: + cur.execute(CREATE_TABLE_SQL) + for row in rows: + cur.execute( + """ + SELECT imei, device_id, serial_number, status + FROM device_imei_mapping + WHERE imei = %s OR device_id = %s + """, + (row["imei"], row["device_id"]), + ) + existing_rows = cur.fetchall() + if len(existing_rows) > 1: + raise ValueError( + f"conflicting mapping: imei={row['imei']} device_id={row['device_id']}" + ) + + if not existing_rows: + cur.execute( + """ + INSERT INTO device_imei_mapping (imei, device_id, serial_number, status) + VALUES (%s, %s, %s, 'pending') + """, + (row["imei"], row["device_id"], row["serial_number"]), + ) + continue + + existing = existing_rows[0] + if existing["imei"] != row["imei"]: + raise ValueError( + f"device_id already mapped to another imei: {row['device_id']}" + ) + + is_activated = str(existing["status"]) == "activated" + changed_identity = ( + existing["device_id"] != row["device_id"] + or existing["serial_number"] != row["serial_number"] + ) + if is_activated and changed_identity: + raise ValueError(f"activated imei cannot be remapped: {row['imei']}") + + cur.execute( + """ + UPDATE device_imei_mapping + SET device_id = %s, + serial_number = %s, + status = IF(status = 'activated', status, 'pending') + WHERE imei = %s + """, + (row["device_id"], row["serial_number"], row["imei"]), + ) + conn.commit() + return len(rows) + except Exception: + conn.rollback() + raise + finally: + conn.close() + + +def main() -> int: + args = parse_args() + csv_path = Path(args.csv_path).expanduser().resolve() + try: + rows = read_rows(csv_path) + if args.dry_run: + print(f"validated {len(rows)} rows from {csv_path}") + return 0 + imported = import_rows(rows) + except Exception as exc: + print(f"Error: {exc}", file=sys.stderr) + return 1 + + print(f"imported {imported} rows into device_imei_mapping") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/talkingq-url/services/device_identity_initializer.py b/talkingq-url/services/device_identity_initializer.py new file mode 100644 index 0000000..8e32481 --- /dev/null +++ b/talkingq-url/services/device_identity_initializer.py @@ -0,0 +1,129 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from sqlalchemy import text + +from services.database_service_base import DatabaseServiceBase +from utils.logger import session_logger + + +class DeviceIdentityInitializationError(ValueError): + pass + + +@dataclass(frozen=True) +class DeviceIdentityInitializationResult: + imei: str + device_id: str + serial_number: str + activated: bool + + +class DeviceIdentityInitializer(DatabaseServiceBase): + def __init__(self): + super().__init__(service_name="device_identity_initializer") + + async def initialize_by_imei(self, imei: str) -> DeviceIdentityInitializationResult: + normalized_imei = str(imei or "").strip() + if not normalized_imei: + raise DeviceIdentityInitializationError("imei is required") + + db_session = await self.get_session() + try: + mapping_result = await db_session.execute( + text( + """ + SELECT imei, device_id, serial_number, status + FROM device_imei_mapping + WHERE imei = :imei + LIMIT 1 + """ + ), + {"imei": normalized_imei}, + ) + mapping = mapping_result.mappings().first() + if mapping is None: + raise DeviceIdentityInitializationError("imei mapping not found") + + device_id = str(mapping["device_id"]).strip() + serial_number = str(mapping["serial_number"]).strip() + if not device_id or not serial_number: + raise DeviceIdentityInitializationError("imei mapping is incomplete") + + device_auth_result = await db_session.execute( + text( + """ + SELECT device_id, serial_number + FROM device_auth + WHERE device_id = :device_id + LIMIT 1 + """ + ), + {"device_id": device_id}, + ) + device_auth = device_auth_result.mappings().first() + if device_auth is None: + await db_session.execute( + text( + """ + INSERT INTO device_auth (device_id, serial_number, is_active) + VALUES (:device_id, :serial_number, 1) + """ + ), + {"device_id": device_id, "serial_number": serial_number}, + ) + elif str(device_auth["serial_number"]) != serial_number: + raise DeviceIdentityInitializationError("device_auth serial_number conflict") + else: + await db_session.execute( + text( + """ + UPDATE device_auth + SET is_active = 1 + WHERE device_id = :device_id + """ + ), + {"device_id": device_id}, + ) + + await db_session.execute( + text( + """ + UPDATE device_imei_mapping + SET status = 'activated', + activated_at = COALESCE(activated_at, CURRENT_TIMESTAMP) + WHERE imei = :imei + """ + ), + {"imei": normalized_imei}, + ) + await db_session.commit() + + session_logger.info( + device_id, + "device_identity_initializer", + f"IMEI identity initialized: imei={normalized_imei}", + ) + return DeviceIdentityInitializationResult( + imei=normalized_imei, + device_id=device_id, + serial_number=serial_number, + activated=str(mapping["status"]) == "activated", + ) + except DeviceIdentityInitializationError: + await db_session.rollback() + raise + except Exception as exc: + await db_session.rollback() + session_logger.error( + normalized_imei, + "device_identity_initializer", + f"initialize identity failed: {exc}", + ) + raise + finally: + await db_session.close() + + +device_identity_initializer = DeviceIdentityInitializer() diff --git a/talkingq-url/services/device_update_manager.py b/talkingq-url/services/device_update_manager.py index 074a77e..d88b511 100644 --- a/talkingq-url/services/device_update_manager.py +++ b/talkingq-url/services/device_update_manager.py @@ -1,6 +1,8 @@ from typing import Dict, Optional import time -from sqlalchemy import select, update, insert, delete +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 @@ -13,7 +15,31 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): self.cache_timestamps = {} # 记录缓存更新时间 self.max_cache_size = 1000 # 最大缓存条目数 - async def create_firmware_update(self, device_id: str, firmware_version: str, update_status: str = "success"): + 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, + "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): """新增固件更新记录""" await self._init_database() db_session = await self.get_session() @@ -29,16 +55,18 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): .values( firmware_version=firmware_version, update_status=update_status, - progress=0.0 + progress=progress ) ) 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=0.0 + progress=progress ) await db_session.execute(insert_stmt) @@ -48,7 +76,7 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): "device_id": device_id, "firmware_version": firmware_version, "update_status": update_status, - "progress": 0.0 + "progress": progress }) return True @@ -79,7 +107,13 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): finally: await db_session.close() - async def update_firmware_update(self, device_id: str, firmware_version: str, update_status: str) -> bool: + async def update_firmware_update( + self, + device_id: str, + firmware_version: str, + update_status: str, + progress: Optional[float] = None, + ) -> bool: """更新设备固件信息(支持部分字段更新)""" await self._init_database() db_session = await self.get_session() @@ -88,6 +122,8 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): "firmware_version": firmware_version, "update_status": update_status } + if progress is not None: + update_values["progress"] = progress query = select(DeviceFirmwareUpdate).where(DeviceFirmwareUpdate.device_id == device_id) result = await db_session.execute(query) @@ -101,10 +137,13 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): ) 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 + update_status=update_status, + progress=progress if progress is not None else 0.0, ) await db_session.execute(insert_stmt) @@ -131,8 +170,12 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): await db_session.execute(update_stmt) await db_session.commit() - if device_id in self.update_cache: - self.update_cache[device_id].progress = progress + 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}%") @@ -184,8 +227,10 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): ) 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" # 默认状态为成功 ) @@ -264,4 +309,10 @@ class DeviceFirmwareUpdateManager(DatabaseServiceBase): if device_id in self.cache_timestamps: del self.cache_timestamps[device_id] -device_firmware_update_manager = DeviceFirmwareUpdateManager() \ No newline at end of file + 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()