Files
archery/ota_manager.py
T

318 lines
10 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.
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
OTA管理器模块
流程:下载ZIP → 解压覆盖项目 → 重启应用程序
"""
import binascii
import hashlib
import threading
import os
import shutil
import requests
import config
from logger_manager import logger_manager
class OTAManager:
"""OTA升级管理器(单例)"""
_instance = None
def __new__(cls):
if cls._instance is None:
cls._instance = super(OTAManager, cls).__new__(cls)
cls._instance._initialized = False
return cls._instance
def __init__(self):
if self._initialized:
return
self._ota_in_progress = 0
self._ota_url = None
self._lock = threading.Lock()
self._initialized = True
@property
def logger(self):
return logger_manager.logger
@property
def ota_in_progress(self):
with self._lock:
return self._ota_in_progress > 0
@property
def update_thread_started(self):
return self._ota_in_progress > 0
@property
def ota_url(self):
return self._ota_url
def _begin_ota(self, url=None):
with self._lock:
self._ota_in_progress += 1
if url:
self._ota_url = url
def _end_ota(self):
with self._lock:
self._ota_in_progress = max(0, self._ota_in_progress - 1)
def _set_ota_url(self, url):
with self._lock:
self._ota_url = url
def _start_update_thread(self):
with self._lock:
if self._ota_in_progress > 0:
return False
self._ota_in_progress += 1
return True
def _stop_update_thread(self):
self._end_ota()
# ==================== 核心方法 ====================
def perform_ota(self, url, progress_callback=None):
"""
完整OTA流程:下载ZIP → 解压覆盖项目
调用方负责重启程序(os.execv)
Args:
url: 固件下载地址
progress_callback: 进度回调 fn(phase, progress)phase="downloading"/"installing"progress=0-100
Returns:
(success: bool, message: str)
"""
if not url:
return False, "missing_url"
self._begin_ota(url)
try:
tmp_path = f"{config.APP_DIR}/ota_tmp.zip"
self.logger.info(f"[OTA] 开始下载: {url}")
if progress_callback:
progress_callback("downloading", 0)
ok, msg = self._download_zip(url, tmp_path, progress_callback)
if not ok:
self.logger.error(f"[OTA] 下载失败: {msg}")
return False, msg
self.logger.info(f"[OTA] 下载完成: {msg}")
if progress_callback:
progress_callback("downloading", 50)
self.logger.info("[OTA] 开始应用更新...")
if progress_callback:
progress_callback("installing", 50)
ok, msg = self._apply_update(tmp_path, progress_callback)
if not ok:
self.logger.error(f"[OTA] 应用更新失败: {msg}")
return False, msg
self.logger.info(f"[OTA] 更新应用成功,共更新 {msg} 个文件")
if progress_callback:
progress_callback("installing", 51)
return True, "success"
except Exception as e:
self.logger.error(f"[OTA] 异常: {e}")
return False, str(e)
finally:
self._end_ota()
def _download_zip(self, url, save_path, progress_callback=None):
"""
下载ZIP文件(流式分块下载,支持进度回调)
Args:
url: 下载地址
save_path: 保存路径
progress_callback: 进度回调 fn(phase, progress)progress=0-80
Returns:
(success: bool, message: str)
"""
try:
response = requests.get(url, timeout=120, stream=True)
response.raise_for_status()
total_size = int(response.headers.get('Content-Length', 0))
chunk_size = 8192
downloaded = 0
md5_hash = hashlib.md5()
with open(save_path, 'wb') as f:
for chunk in response.iter_content(chunk_size=chunk_size):
if chunk:
f.write(chunk)
md5_hash.update(chunk)
downloaded += len(chunk)
if progress_callback and total_size > 0:
percent = min(int(downloaded / total_size * 50), 49)
progress_callback("downloading", percent)
try:
os.sync()
except:
pass
md5_b64_expected = None
if 'Content-Md5' in response.headers:
md5_b64_expected = response.headers['Content-Md5'].strip()
if md5_b64_expected:
md5_b64_got = binascii.b2a_base64(md5_hash.digest()).decode().strip()
if md5_b64_got != md5_b64_expected:
return False, f"MD5校验失败"
self.logger.info("[OTA] MD5校验通过")
return True, f"size={downloaded}"
except requests.exceptions.RequestException as e:
return False, f"网络错误: {e}"
except OSError as e:
return False, f"写入错误: {e}"
def _apply_update(self, zip_path, progress_callback=None):
"""
解压ZIP到临时目录,重启后由主程序移动到实际目录
Returns:
(success: bool, message: str)
"""
if not os.path.exists(zip_path):
return False, f"文件不存在: {zip_path}"
try:
with open(zip_path, "rb") as f:
header = f.read(4)
if header[:2] != b'PK':
return False, f"不是ZIP文件: {header.hex()}"
except Exception as e:
return False, f"读取ZIP失败: {e}"
staging_dir = f"{config.APP_DIR}/ota_staging"
try:
os.makedirs(staging_dir, exist_ok=True)
except:
pass
try:
self.logger.info(f"[OTA] 开始解压: {zip_path} -> {staging_dir}")
ret = os.system(f"unzip -q -o '{zip_path}' -d '{staging_dir}' 2>&1")
if ret != 0:
return False, f"解压失败: exit={ret}"
self.logger.info("[OTA] 解压完成")
except Exception as e:
return False, f"解压异常: {e}"
file_count = 0
for _, _, files in os.walk(staging_dir):
file_count += len(files)
if file_count == 0:
return False, "ZIP中无文件"
try:
os.sync()
except:
pass
try:
os.remove(zip_path)
except:
pass
self.logger.info(f"[OTA] 已解压 {file_count} 个文件到临时目录,重启后生效")
return True, file_count
def restore_from_backup(self, backup_dir_path=None):
"""
从备份目录恢复所有文件到应用目录
Args:
backup_dir_path: 备份目录路径,None则自动查找最新备份
Returns:
bool: 是否成功恢复
"""
backup_base = config.BACKUP_BASE
try:
if backup_dir_path is None:
if not os.path.exists(backup_base):
self.logger.error(f"[RESTORE] 备份目录不存在: {backup_base}")
return False
backup_dirs = []
for item in os.listdir(backup_base):
if item == ".counter":
continue
item_path = os.path.join(backup_base, item)
if os.path.isdir(item_path) and item.startswith("backup_"):
try:
dir_num = int(item.replace("backup_", ""))
backup_dirs.append((item, dir_num))
except:
pass
if not backup_dirs:
self.logger.error("[RESTORE] 没有找到备份目录")
return False
backup_dirs.sort(key=lambda x: x[1], reverse=True)
backup_dir_path = os.path.join(backup_base, backup_dirs[0][0])
if not os.path.exists(backup_dir_path):
self.logger.error(f"[RESTORE] 备份目录不存在: {backup_dir_path}")
return False
self.logger.info(f"[RESTORE] 开始从备份恢复: {backup_dir_path}")
restored_files = []
for root, dirs, files in os.walk(backup_dir_path):
for f in files:
src = os.path.join(root, f)
rel = os.path.relpath(src, backup_dir_path)
dest = os.path.join(config.APP_DIR, rel)
dest_dir = os.path.dirname(dest)
if dest_dir:
os.makedirs(dest_dir, exist_ok=True)
try:
shutil.copy2(src, dest)
restored_files.append(rel)
except Exception as e:
self.logger.error(f"[RESTORE] 恢复 {rel} 失败: {e}")
if restored_files:
self.logger.info(f"[RESTORE] 成功恢复 {len(restored_files)} 个文件")
return True
else:
self.logger.info("[RESTORE] 没有文件被恢复")
return False
except Exception as e:
self.logger.error(f"[RESTORE] 恢复过程出错: {e}")
return False
# 全局单例
ota_manager = OTAManager()
# ==================== 向后兼容接口 ====================
def apply_ota_and_reboot(ota_url=None, downloaded_file=None):
return ota_manager.perform_ota(ota_url)
def direct_ota_download_via_4g(ota_url):
return ota_manager.perform_ota(ota_url)
def handle_wifi_and_update(ssid, password, ota_url):
return ota_manager.perform_ota(ota_url)
def restore_from_backup(backup_dir_path=None):
return ota_manager.restore_from_backup(backup_dir_path)