from __future__ import annotations import hashlib import ipaddress import json from pathlib import Path 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}$") class ImageConfigError(ValueError): pass def validate_config(value: Any) -> dict[str, Any]: if not isinstance(value, dict) or set(value) != { "schema_version", "product", "software_version", "account", "wifi", "ipv4", }: raise ImageConfigError("image configuration fields are invalid") if value["schema_version"] != 1 or value["product"] != PRODUCT_ID: raise ImageConfigError("image configuration product or schema is unsupported") if not isinstance(value["software_version"], str) or not re.fullmatch( r"(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)\.(?:0|[1-9][0-9]*)", value["software_version"], ): raise ImageConfigError("software_version must use MAJOR.MINOR.PATCH") account = value["account"] if not isinstance(account, dict) or set(account) != {"username", "password"}: raise ImageConfigError("account configuration is invalid") if not isinstance(account["username"], str) or not _USERNAME.fullmatch(account["username"]): raise ImageConfigError("account username is not a supported Linux username") _validate_secret(account["password"], "account password", minimum=8, maximum=128) wifi = value["wifi"] if not isinstance(wifi, dict) or set(wifi) != {"ssid", "password"}: raise ImageConfigError("Wi-Fi configuration is invalid") _validate_text(wifi["ssid"], "Wi-Fi SSID", minimum=1, maximum_bytes=32) _validate_secret(wifi["password"], "Wi-Fi password", minimum=8, maximum=63) ipv4 = value["ipv4"] if not isinstance(ipv4, dict) or set(ipv4) != { "mode", "address", "prefix", "gateway", "dns", }: raise ImageConfigError("IPv4 configuration is invalid") if ipv4["mode"] not in {"dhcp", "static"}: raise ImageConfigError("IPv4 mode must be dhcp or static") if ipv4["mode"] == "dhcp": if ipv4 != {"mode": "dhcp", "address": "", "prefix": 0, "gateway": "", "dns": []}: raise ImageConfigError("DHCP configuration must not include static fields") else: try: address = ipaddress.IPv4Address(ipv4["address"]) gateway = ipaddress.IPv4Address(ipv4["gateway"]) except (ipaddress.AddressValueError, TypeError) as exc: raise ImageConfigError("static IPv4 address or gateway is invalid") from exc prefix = ipv4["prefix"] if type(prefix) is not int or not 1 <= prefix <= 32: raise ImageConfigError("static IPv4 prefix must be between 1 and 32") network = ipaddress.IPv4Network(f"{address}/{prefix}", strict=False) if gateway not in network: raise ImageConfigError("static IPv4 gateway must be in the same subnet") if not isinstance(ipv4["dns"], list) or not 1 <= len(ipv4["dns"]) <= 4: raise ImageConfigError("static IPv4 DNS must contain one to four addresses") try: ipv4["dns"] = [str(ipaddress.IPv4Address(item)) for item in ipv4["dns"]] except (ipaddress.AddressValueError, TypeError) as exc: raise ImageConfigError("static IPv4 DNS contains an invalid address") from exc return value def _validate_text(value: Any, label: str, *, minimum: int, maximum_bytes: int) -> None: if not isinstance(value, str) or len(value) < minimum or len(value.encode("utf-8")) > maximum_bytes: raise ImageConfigError(f"{label} length is invalid") if any(character in value for character in ("\0", "\r", "\n")): raise ImageConfigError(f"{label} contains a forbidden character") def _validate_secret(value: Any, label: str, *, minimum: int, maximum: int) -> None: if not isinstance(value, str) or not minimum <= len(value) <= maximum: raise ImageConfigError(f"{label} length is invalid") if any(character in value for character in ("\0", "\r", "\n")): raise ImageConfigError(f"{label} contains a forbidden character") def _slot_offset(index: int) -> int: return HEADER_BYTES + index * SLOT_BYTES def _encode_slot(config: dict[str, Any], generation: int) -> bytes: payload = (json.dumps(validate_config(config), ensure_ascii=False, sort_keys=True, separators=(",", ":")) + "\n").encode("utf-8") maximum = SLOT_BYTES - _SLOT_HEADER.size if len(payload) > maximum: raise ImageConfigError("image configuration payload is too large") 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: dict[str, 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("image configuration file size is invalid") magic, version, declared_size = _HEADER.unpack_from(data, 0) if magic != FILE_MAGIC or version != FORMAT_VERSION or declared_size != CONFIG_FILE_BYTES: raise ImageConfigError("image configuration header is invalid") 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 = data[offset + _SLOT_HEADER.size : offset + _SLOT_HEADER.size + payload_bytes] if hashlib.sha256(payload).digest() != digest: continue try: document = validate_config(json.loads(payload.decode("utf-8"))) except (UnicodeError, json.JSONDecodeError, ImageConfigError): continue valid.append((generation, index, document)) if not valid: raise ImageConfigError("image configuration has no valid slot") generation, index, document = max(valid, key=lambda item: (item[0], item[1])) return document, generation, index def update_config_file(data: bytes, config: dict[str, Any]) -> bytes: _current, generation, active = read_config_file(data) target = 1 - active result = bytearray(data) offset = _slot_offset(target) result[offset : offset + SLOT_BYTES] = _encode_slot(config, generation + 1) return bytes(result) def read_config_path(path: Path) -> dict[str, Any]: document, _generation, _slot = read_config_file(Path(path).read_bytes()) return document