201 lines
8.1 KiB
Python
201 lines
8.1 KiB
Python
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import ipaddress
|
|
import json
|
|
import re
|
|
import struct
|
|
from typing import Any
|
|
|
|
|
|
PRODUCT_ID = "matrix-screen-controller-walnutpi"
|
|
CONFIG_FILE_BYTES = 64 * 1024
|
|
HEADER_BYTES = 4096
|
|
SLOT_BYTES = (CONFIG_FILE_BYTES - HEADER_BYTES) // 2
|
|
FILE_MAGIC = b"MSCCFG2\0"
|
|
SLOT_MAGIC = b"MSCSLOT\0"
|
|
FORMAT_VERSION = 1
|
|
_HEADER = struct.Struct("<8sII")
|
|
_SLOT_HEADER = struct.Struct("<8sQI32s")
|
|
_USERNAME = re.compile(r"[a-z_][a-z0-9_-]{0,31}", re.ASCII)
|
|
_VERSION = re.compile(r"(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)", re.ASCII)
|
|
|
|
|
|
class ImageConfigError(ValueError):
|
|
pass
|
|
|
|
|
|
def _exact_object(value: Any, fields: set[str], label: str) -> dict[str, Any]:
|
|
if not isinstance(value, dict) or set(value) != fields:
|
|
raise ImageConfigError(f"{label}字段无效")
|
|
return value
|
|
|
|
|
|
def _utf16_units(value: str) -> int:
|
|
# .NET String.Length counts UTF-16 code units; retain 1.0.1 validation semantics.
|
|
return len(value.encode("utf-16-le")) // 2
|
|
|
|
|
|
def _validate_secret(value: Any, label: str, minimum: int, maximum: int) -> str:
|
|
if not isinstance(value, str) or not minimum <= _utf16_units(value) <= maximum:
|
|
raise ImageConfigError(f"{label}长度或字符无效")
|
|
if any(character in value for character in ("\0", "\r", "\n")):
|
|
raise ImageConfigError(f"{label}长度或字符无效")
|
|
return value
|
|
|
|
|
|
def _validate_text(value: Any, label: str, minimum: int, maximum_bytes: int) -> str:
|
|
if not isinstance(value, str) or _utf16_units(value) < minimum or len(value.encode("utf-8")) > maximum_bytes:
|
|
raise ImageConfigError(f"{label}长度或字符无效")
|
|
if any(character in value for character in ("\0", "\r", "\n")):
|
|
raise ImageConfigError(f"{label}长度或字符无效")
|
|
return value
|
|
|
|
|
|
def validate_config(value: Any) -> dict[str, Any]:
|
|
source = _exact_object(
|
|
value,
|
|
{"schema_version", "product", "software_version", "account", "wifi", "ipv4"},
|
|
"镜像配置",
|
|
)
|
|
schema = source["schema_version"]
|
|
if type(schema) is not int or schema != FORMAT_VERSION or source["product"] != PRODUCT_ID:
|
|
raise ImageConfigError("镜像产品或配置版本不受支持")
|
|
software_version = source["software_version"]
|
|
if not isinstance(software_version, str) or _VERSION.fullmatch(software_version) is None:
|
|
raise ImageConfigError("软件版本格式无效")
|
|
|
|
account = _exact_object(source["account"], {"username", "password"}, "账户配置")
|
|
username = account["username"]
|
|
if not isinstance(username, str) or _USERNAME.fullmatch(username) is None:
|
|
raise ImageConfigError("用户名只能使用小写字母、数字、下划线和连字符,且最长 32 个字符")
|
|
account_password = _validate_secret(account["password"], "账户密码", 8, 128)
|
|
|
|
wifi = _exact_object(source["wifi"], {"ssid", "password"}, "Wi-Fi 配置")
|
|
ssid = _validate_text(wifi["ssid"], "Wi-Fi 名称", 1, 32)
|
|
wifi_password = _validate_secret(wifi["password"], "Wi-Fi 密码", 8, 63)
|
|
|
|
ipv4 = _exact_object(source["ipv4"], {"mode", "address", "prefix", "gateway", "dns"}, "IPv4 配置")
|
|
mode = ipv4["mode"]
|
|
if mode == "dhcp":
|
|
normalized_ipv4: dict[str, Any] = {"mode": "dhcp", "address": "", "prefix": 0, "gateway": "", "dns": []}
|
|
elif mode == "static":
|
|
try:
|
|
address = ipaddress.IPv4Address(ipv4["address"])
|
|
gateway = ipaddress.IPv4Address(ipv4["gateway"])
|
|
except (ipaddress.AddressValueError, TypeError) as exception:
|
|
raise ImageConfigError("静态 IPv4 地址或网关无效") from exception
|
|
prefix = ipv4["prefix"]
|
|
if type(prefix) is not int or not 1 <= prefix <= 32:
|
|
raise ImageConfigError("IPv4 前缀必须是 1 到 32")
|
|
if gateway not in ipaddress.IPv4Network(f"{address}/{prefix}", strict=False):
|
|
raise ImageConfigError("静态 IPv4 网关必须与地址处于同一子网")
|
|
dns_source = ipv4["dns"]
|
|
if not isinstance(dns_source, list) or not 1 <= len(dns_source) <= 4:
|
|
raise ImageConfigError("DNS 必须包含 1 到 4 个 IPv4 地址")
|
|
try:
|
|
dns = [str(ipaddress.IPv4Address(item)) for item in dns_source]
|
|
except (ipaddress.AddressValueError, TypeError) as exception:
|
|
raise ImageConfigError("DNS 必须包含 1 到 4 个 IPv4 地址") from exception
|
|
normalized_ipv4 = {
|
|
"mode": "static",
|
|
"address": str(ipv4["address"]),
|
|
"prefix": prefix,
|
|
"gateway": str(ipv4["gateway"]),
|
|
"dns": dns,
|
|
}
|
|
else:
|
|
raise ImageConfigError("IPv4 模式无效")
|
|
|
|
return {
|
|
"schema_version": FORMAT_VERSION,
|
|
"product": PRODUCT_ID,
|
|
"software_version": software_version,
|
|
"account": {"username": username, "password": account_password},
|
|
"wifi": {"ssid": ssid, "password": wifi_password},
|
|
"ipv4": normalized_ipv4,
|
|
}
|
|
|
|
|
|
def serialize_config(value: Any) -> str:
|
|
return json.dumps(validate_config(value), ensure_ascii=False, separators=(",", ":"))
|
|
|
|
|
|
def deserialize_config(text: str) -> dict[str, Any]:
|
|
try:
|
|
value = json.loads(text)
|
|
except (UnicodeError, json.JSONDecodeError) as exception:
|
|
raise ImageConfigError("配置 JSON 无效") from exception
|
|
return validate_config(value)
|
|
|
|
|
|
def _slot_offset(index: int) -> int:
|
|
return HEADER_BYTES + index * SLOT_BYTES
|
|
|
|
|
|
def _encode_slot(config: Any, generation: int) -> bytes:
|
|
payload = (serialize_config(config) + "\n").encode("utf-8")
|
|
if len(payload) > SLOT_BYTES - _SLOT_HEADER.size:
|
|
raise ImageConfigError("配置内容过大")
|
|
header = _SLOT_HEADER.pack(SLOT_MAGIC, generation, len(payload), hashlib.sha256(payload).digest())
|
|
return header + payload + bytes(SLOT_BYTES - len(header) - len(payload))
|
|
|
|
|
|
def create_config_file(config: Any) -> bytes:
|
|
result = bytearray(CONFIG_FILE_BYTES)
|
|
_HEADER.pack_into(result, 0, FILE_MAGIC, FORMAT_VERSION, CONFIG_FILE_BYTES)
|
|
result[_slot_offset(0) : _slot_offset(0) + SLOT_BYTES] = _encode_slot(config, 1)
|
|
return bytes(result)
|
|
|
|
|
|
def read_config_file(data: bytes) -> tuple[dict[str, Any], int, int]:
|
|
if len(data) != CONFIG_FILE_BYTES:
|
|
raise ImageConfigError("镜像配置文件大小无效")
|
|
try:
|
|
magic, version, declared_size = _HEADER.unpack_from(data, 0)
|
|
except struct.error as exception:
|
|
raise ImageConfigError("镜像配置区头部无效") from exception
|
|
if magic != FILE_MAGIC or version != FORMAT_VERSION or declared_size != CONFIG_FILE_BYTES:
|
|
raise ImageConfigError("镜像配置区头部无效")
|
|
|
|
valid: list[tuple[int, int, dict[str, Any]]] = []
|
|
for index in range(2):
|
|
offset = _slot_offset(index)
|
|
slot_magic, generation, payload_bytes, digest = _SLOT_HEADER.unpack_from(data, offset)
|
|
if slot_magic != SLOT_MAGIC or payload_bytes <= 0 or payload_bytes > SLOT_BYTES - _SLOT_HEADER.size:
|
|
continue
|
|
payload_start = offset + _SLOT_HEADER.size
|
|
payload = data[payload_start : payload_start + payload_bytes]
|
|
if hashlib.sha256(payload).digest() != digest:
|
|
continue
|
|
try:
|
|
document = deserialize_config(payload.decode("utf-8"))
|
|
except (UnicodeError, ImageConfigError):
|
|
continue
|
|
valid.append((generation, index, document))
|
|
if not valid:
|
|
raise ImageConfigError("镜像配置区没有可恢复的有效副本")
|
|
generation, index, document = max(valid, key=lambda item: (item[0], item[1]))
|
|
return document, generation, index
|
|
|
|
|
|
def encode_update(data: bytes, config: Any) -> tuple[bytes, int]:
|
|
_current, generation, active = read_config_file(data)
|
|
target = 1 - active
|
|
return _encode_slot(config, generation + 1), target
|
|
|
|
|
|
def update_config_file(data: bytes, config: Any) -> bytes:
|
|
slot_data, target = encode_update(data, config)
|
|
result = bytearray(data)
|
|
offset = _slot_offset(target)
|
|
result[offset : offset + SLOT_BYTES] = slot_data
|
|
return bytes(result)
|
|
|
|
|
|
def slot_offset(index: int) -> int:
|
|
if index not in (0, 1):
|
|
raise ValueError("配置槽索引无效")
|
|
return _slot_offset(index)
|
|
|