合入小程序修改,新增告警功能

This commit is contained in:
HycJack
2026-05-09 09:46:06 +08:00
parent 46cc5272f3
commit be4d276557
10 changed files with 800 additions and 28 deletions

View File

@@ -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 = (

View File

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

View File

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

View File

@@ -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:

View File

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

View File

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

View File

@@ -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")

View File

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

View File

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

View File

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