Files

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)