From 33a385cfa8b4f4bae0a7155326d1a2bee8b0a506 Mon Sep 17 00:00:00 2001 From: rui chen Date: Wed, 17 Dec 2025 16:26:35 +0800 Subject: [PATCH] add basic OTA support for single server deployment Committer: rxchen --- main/xiaozhi-server/core/api/ota_handler.py | 219 ++++++++++++++++++-- main/xiaozhi-server/core/http_server.py | 5 +- 2 files changed, 209 insertions(+), 15 deletions(-) diff --git a/main/xiaozhi-server/core/api/ota_handler.py b/main/xiaozhi-server/core/api/ota_handler.py index b6c88dff..a521332e 100644 --- a/main/xiaozhi-server/core/api/ota_handler.py +++ b/main/xiaozhi-server/core/api/ota_handler.py @@ -3,6 +3,10 @@ import time import base64 import hashlib import hmac +import os +import re +import glob +from typing import Dict, List, Tuple from aiohttp import web from core.auth import AuthManager @@ -12,6 +16,33 @@ from core.api.base_handler import BaseHandler TAG = __name__ +def _safe_basename(filename: str) -> str: + # Prevent directory traversal + return os.path.basename(filename) + + +def _parse_version(ver: str) -> Tuple[int, ...]: + # conservative parser: split by non-digit, keep numeric parts + parts = re.findall(r"\d+", ver) + return tuple(int(p) for p in parts) if parts else (0,) + + +def _is_higher_version(a: str, b: str) -> bool: + """Return True if version string a > b (semver-like numeric compare).""" + ta = _parse_version(a) + tb = _parse_version(b) + # compare tuple lexicographically, but allow different lengths + maxlen = max(len(ta), len(tb)) + for i in range(maxlen): + ai = ta[i] if i < len(ta) else 0 + bi = tb[i] if i < len(tb) else 0 + if ai > bi: + return True + if ai < bi: + return False + return False + + class OTAHandler(BaseHandler): def __init__(self, config: dict): super().__init__(config) @@ -23,6 +54,46 @@ class OTAHandler(BaseHandler): expire_seconds = auth_config.get("expire_seconds") self.auth = AuthManager(secret_key=secret_key, expire_seconds=expire_seconds) + # firmware storage + self.bin_dir = os.path.join(os.getcwd(), "data", "bin") + # cache structure: { 'updated_at': timestamp, 'ttl': seconds, 'files_by_model': { model: [(version, filename), ...] } } + self._bin_cache: Dict = {"updated_at": 0, "ttl": config.get("firmware_cache_ttl", 30), "files_by_model": {}} + + def _refresh_bin_cache_if_needed(self): + now = int(time.time()) + ttl = int(self._bin_cache.get("ttl", 30)) + if now - int(self._bin_cache.get("updated_at", 0)) < ttl and self._bin_cache.get("files_by_model"): + return + + files_by_model: Dict[str, List[Tuple[str, str]]] = {} + try: + if not os.path.isdir(self.bin_dir): + os.makedirs(self.bin_dir, exist_ok=True) + + # match files like model_1.2.3.bin (allow dots, dashes, underscores in model and version) + pattern = os.path.join(self.bin_dir, "*.bin") + for path in glob.glob(pattern): + fname = os.path.basename(path) + # filename format: {model}_{version}.bin + m = re.match(r"^(.+?)_([0-9][A-Za-z0-9\.\-_]*)\.bin$", fname) + if not m: + # skip files not conforming to naming rule + continue + model = m.group(1) + version = m.group(2) + files_by_model.setdefault(model, []).append((version, fname)) + + # sort versions for each model descending + for model, items in files_by_model.items(): + items.sort(key=lambda it: _parse_version(it[0]), reverse=True) + + self._bin_cache["files_by_model"] = files_by_model + self._bin_cache["updated_at"] = now + self.logger.bind(tag=TAG).info(f"Firmware cache refreshed: {len(files_by_model)} models") + except Exception as e: + self.logger.bind(tag=TAG).error(f"刷新固件缓存失败: {e}") + # keep previous cache if any + def generate_password_signature(self, content: str, secret_key: str) -> str: """生成MQTT密码签名 @@ -62,7 +133,14 @@ class OTAHandler(BaseHandler): return f"ws://{local_ip}:{port}/xiaozhi/v1/" async def handle_post(self, request): - """处理 OTA POST 请求""" + """处理 OTA POST 请求 + + This handler will: + - read device id/client id (as before) + - attempt to determine device model and current firmware version (prefer headers, fallback to body) + - check data/bin for newer firmware for that model + - if found a newer firmware, set firmware.url to the download endpoint + """ try: data = await request.text() self.logger.bind(tag=TAG).debug(f"OTA请求方法: {request.method}") @@ -81,11 +159,54 @@ class OTAHandler(BaseHandler): else: raise Exception("OTA请求ClientID为空") - data_json = json.loads(data) + data_json = {} + try: + data_json = json.loads(data) if data else {} + self.logger.bind(tag=TAG).info(f"data json:{data_json}") + except Exception: + data_json = {} server_config = self.config["server"] - port = int(server_config.get("port", 8000)) + # Distinguish ports: + # - websocket_port is used to construct websocket URL (server["port"]) + # - http_port is used to construct OTA download URLs (server["http_port"]) + websocket_port = int(server_config.get("port", 8000)) + http_port = int(server_config.get("http_port", 8003)) local_ip = get_local_ip() + ota_addr = server_config.get("ota_addr", "") + + # Determine device model (prefer headers) + device_model = "" + # header candidates + for h in ("device-model", "device_model", "model"): + if h in request.headers: + device_model = request.headers.get(h, "").strip() + break + # body fallback + if not device_model: + try: + if "board" in data_json and isinstance(data_json["board"], dict): + device_model = data_json["board"].get("type", "") + elif "model" in data_json: + device_model = data_json.get("model", "") + except Exception: + device_model = "" + if not device_model: + device_model = "default" + + # Determine device current version (prefer headers) + device_version = "" + for h in ("device-version", "device_version", "firmware-version", "app-version", "application-version"): + if h in request.headers: + device_version = request.headers.get(h, "").strip() + break + if not device_version: + try: + device_version = data_json.get("application", {}).get("version", "") + except Exception: + device_version = "" + if not device_version: + device_version = "0.0.0" return_json = { "server_time": { @@ -93,21 +214,17 @@ class OTAHandler(BaseHandler): "timezone_offset": server_config.get("timezone_offset", 8) * 60, }, "firmware": { - "version": data_json["application"].get("version", "1.0.0"), + "version": device_version, "url": "", }, } + # existing mqtt/websocket logic (unchanged) mqtt_gateway_endpoint = server_config.get("mqtt_gateway") if mqtt_gateway_endpoint: # 如果配置了非空字符串 - # 尝试从请求数据中获取设备型号 - device_model = "default" + # 尝试从请求数据中获取设备型号(已解析 above) try: - if "device" in data_json and isinstance(data_json["device"], dict): - device_model = data_json["device"].get("model", "default") - elif "model" in data_json: - device_model = data_json["model"] group_id = f"GID_{device_model}".replace(":", "_").replace(" ", "_") except Exception as e: self.logger.bind(tag=TAG).error(f"获取设备型号失败: {e}") @@ -159,20 +276,51 @@ class OTAHandler(BaseHandler): token = self.auth.generate_token(client_id, device_id) else: token = self.auth.generate_token(client_id, device_id) + # NOTE: use websocket_port here return_json["websocket"] = { - "url": self._get_websocket_url(local_ip, port), + "url": self._get_websocket_url(local_ip, websocket_port), "token": token, } self.logger.bind(tag=TAG).info( f"未配置MQTT网关,为设备 {device_id} 下发WebSocket配置" ) - self.logger.bind(tag=TAG).info(f"{return_json}") + + # Now check firmware files for updates + try: + self._refresh_bin_cache_if_needed() + files_by_model = self._bin_cache.get("files_by_model", {}) + candidates = files_by_model.get(device_model, []) + + self.logger.bind(tag=TAG).info(f"查找型号 {device_model} 的固件,找到 {len(candidates)} 个候选") + + chosen_url = "" + chosen_version = device_version + + # candidates are sorted descending by version + for ver, fname in candidates: + if _is_higher_version(ver, device_version): + # build download url (only allow download via our download endpoint) + chosen_version = ver + # use local_ip and http_port to construct url + chosen_url = f"http://{ota_addr}:{http_port}/xiaozhi/ota/download/{fname}" + break + + if chosen_url: + return_json["firmware"]["version"] = chosen_version + return_json["firmware"]["url"] = chosen_url + self.logger.bind(tag=TAG).info(f"为设备 {device_id} 下发固件 {chosen_version} -> {chosen_url}") + else: + self.logger.bind(tag=TAG).info(f"设备 {device_id} 固件已是最新: {device_version}") + + except Exception as e: + self.logger.bind(tag=TAG).error(f"检查固件版本时出错: {e}") response = web.Response( text=json.dumps(return_json, separators=(",", ":")), content_type="application/json", ) except Exception as e: + self.logger.bind(tag=TAG).error(f"OTA POST处理异常: {e}") return_json = {"success": False, "message": "request error."} response = web.Response( text=json.dumps(return_json, separators=(",", ":")), @@ -187,8 +335,9 @@ class OTAHandler(BaseHandler): try: server_config = self.config["server"] local_ip = get_local_ip() - port = int(server_config.get("port", 8000)) - websocket_url = self._get_websocket_url(local_ip, port) + # use websocket port for websocket URL + websocket_port = int(server_config.get("port", 8000)) + websocket_url = self._get_websocket_url(local_ip, websocket_port) message = f"OTA接口运行正常,向设备发送的websocket地址是:{websocket_url}" response = web.Response(text=message, content_type="text/plain") except Exception as e: @@ -197,3 +346,45 @@ class OTAHandler(BaseHandler): finally: self._add_cors_headers(response) return response + + async def handle_download(self, request): + """ + 下载固件接口 + URL: /xiaozhi/ota/download/{filename} + - 只允许下载 data/bin 目录下的 .bin 文件 + - filename 必须是 basename 且匹配安全的模式 + """ + try: + fname = request.match_info.get("filename", "") + if not fname: + raise web.HTTPBadRequest(text="filename required") + + # sanitize + fname = _safe_basename(fname) + # pattern: allow letters, numbers, dot, underscore, dash + if not re.match(r"^[A-Za-z0-9\.\-_]+\.bin$", fname): + raise web.HTTPBadRequest(text="invalid filename") + + file_path = os.path.join(self.bin_dir, fname) + # ensure realpath is under bin_dir + file_real = os.path.realpath(file_path) + bin_dir_real = os.path.realpath(self.bin_dir) + if not file_real.startswith(bin_dir_real + os.sep) and file_real != bin_dir_real: + raise web.HTTPForbidden(text="forbidden") + + if not os.path.isfile(file_real): + raise web.HTTPNotFound(text="file not found") + + # use FileResponse to stream file + resp = web.FileResponse(path=file_real) + except web.HTTPError as e: + resp = e + except Exception as e: + self.logger.bind(tag=TAG).error(f"固件下载异常: {e}") + resp = web.Response(text="download error", status=500) + finally: + try: + self._add_cors_headers(resp) + except Exception: + pass + return resp diff --git a/main/xiaozhi-server/core/http_server.py b/main/xiaozhi-server/core/http_server.py index edbdf1fe..ecc80efb 100644 --- a/main/xiaozhi-server/core/http_server.py +++ b/main/xiaozhi-server/core/http_server.py @@ -48,6 +48,9 @@ class SimpleHttpServer: web.get("/xiaozhi/ota/", self.ota_handler.handle_get), web.post("/xiaozhi/ota/", self.ota_handler.handle_post), web.options("/xiaozhi/ota/", self.ota_handler.handle_post), + # 下载接口,仅提供 data/bin/*.bin 下载 + web.get("/xiaozhi/ota/download/{filename}", self.ota_handler.handle_download), + web.options("/xiaozhi/ota/download/{filename}", self.ota_handler.handle_download), ] ) # 添加路由 @@ -67,4 +70,4 @@ class SimpleHttpServer: # 保持服务运行 while True: - await asyncio.sleep(3600) # 每隔 1 小时检查一次 + await asyncio.sleep(3600) # 每隔 1 小时检查一次 \ No newline at end of file