#!/usr/bin/env python3
"""Verify the JetHome device signature stored in an EEPROM dump.

Usage:
    python3 verify_eeprom_signature.py DUMP [--key PUBLIC_KEY.pem]

DUMP holds the first 540 bytes of the CPU module EEPROM - the board
header, the header of the first file and the device.id record - as raw
binary or as a hex string. Without --key the JetHome public key matching
the record's signature_version is taken from the script's directory.

Requirements: pip install jeefs cryptography

Exit status: 0 - signature valid, 1 - signature invalid, 2 - the check
could not run: the board header or the device.id record is missing,
damaged or unsupported, the filesystem version is unknown, the record
carries no signature, or the dump or the key cannot be read.
"""

import argparse
import binascii
import struct
import sys
from datetime import datetime, timezone
from pathlib import Path

from cryptography.exceptions import InvalidSignature, UnsupportedAlgorithm
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric import ec
from cryptography.hazmat.primitives.asymmetric.utils import encode_dss_signature
from cryptography.hazmat.primitives.serialization import load_pem_public_key
from jeefs import (
    DEVICE_ID_FILENAME,
    EEPROM_MAGIC,
    DeviceIdentityV1,
    EEPROMHeaderV3,
    EEPROMHeaderV4,
    SignatureAlgorithm,
    detect_version,
)

HEADER_CLASSES = {3: EEPROMHeaderV3, 4: EEPROMHeaderV4}
HEADER_SIZE = 256
FS_VERSION = 1  # the current filesystem version
# File header: name, dataSize, crc32, nextFileAddress, headerCrc32.
FILE_HEADER = struct.Struct("<16sHIHI")
RECORD_SIZE = 256
# device.id is always the first file, right after the board header.
RECORD_OFFSET = HEADER_SIZE + FILE_HEADER.size
DUMP_SIZE = RECORD_OFFSET + RECORD_SIZE

# signature_version -> (curve, default public key file)
ALGORITHMS = {
    SignatureAlgorithm.SECP192R1: (ec.SECP192R1, "secp192r1.pem"),
    SignatureAlgorithm.SECP256R1: (ec.SECP256R1, "secp256r1.pem"),
}


def show(name: str, value: str) -> None:
    print(f"{name:<14}{value}")


def crc32(data: bytes) -> int:
    return binascii.crc32(data) & 0xFFFFFFFF


def load_dump(path: Path) -> bytes:
    data = path.read_bytes()
    try:
        return bytes.fromhex(data.decode("ascii"))
    except ValueError:
        return data  # not a hex string: a raw binary dump


def device_record(data: bytes) -> bytes:
    """Return the device.id record from the first file slot.

    Raises ValueError with the reason when there is no valid record.
    """
    if len(data) < DUMP_SIZE:
        raise ValueError(f"the dump is too short: read at least {DUMP_SIZE} bytes")
    raw = data[HEADER_SIZE:RECORD_OFFSET]
    if raw[0] in (0x00, 0xFF):
        raise ValueError("there are no files: the device is not signed")
    name, size, data_crc, _next, header_crc = FILE_HEADER.unpack(raw)
    if crc32(raw[:-4]) != header_crc:
        raise ValueError("file header CRC32 mismatch")
    if name.split(b"\0")[0] != DEVICE_ID_FILENAME.encode() or size != RECORD_SIZE:
        raise ValueError(f"the first file is not {DEVICE_ID_FILENAME}: the device is not signed")
    record = data[RECORD_OFFSET:DUMP_SIZE]
    if crc32(record) != data_crc:
        raise ValueError("file data CRC32 mismatch")
    return record


def signed_payload(header: EEPROMHeaderV3) -> bytes:
    """Rebuild the signed string: "<cpuid>:<mac>:<usid>".

    The values come from the board header. The mac field is written as
    12 upper-case hex digits without separators; cpuid and usid are
    taken verbatim.
    """
    mac = header.mac.replace(":", "")
    return f"{header.cpuid}:{mac}:{header.usid}".encode("utf-8")


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    parser.add_argument("dump", type=Path, help="EEPROM dump, binary or hex")
    parser.add_argument("--key", type=Path, help="public key in PEM format")
    args = parser.parse_args()

    try:
        data = load_dump(args.dump)
    except OSError as err:
        print(f"{args.dump}: {err.strerror}")
        return 2

    version = detect_version(data)
    if version is None:
        if data[:8] != EEPROM_MAGIC:
            print("No JetHome header: the magic does not match")
        elif len(data) < 12:
            # detect_version() needs the magic, the version byte and the
            # three bytes after it.
            print(f"Dump too short: {len(data)} bytes")
        else:
            print(f"Unsupported header version: {data[8]}")
        return 2
    if version not in HEADER_CLASSES:
        print(f"Header v{version} is not supported by this script")
        return 2
    if not EEPROMHeaderV3.verify_crc_static(data):
        print("Board header CRC32 mismatch: the dump is damaged or incomplete")
        return 2
    try:
        header = HEADER_CLASSES[version].from_bytes(data)
    except ValueError as err:
        print(f"Unsupported board header: {err}")
        return 2
    show("header", f"v{version}")
    show("board", f"{header.boardname} {header.boardversion}")
    show(header.SERIAL_LABEL, header.serial)
    show("cpuid", header.cpuid)
    show("mac", header.mac)
    show("usid", header.usid)

    if header.fs_version > FS_VERSION:
        print(f"Unsupported filesystem version: {header.fs_version}")
        return 2
    try:
        raw = device_record(data)
        record = DeviceIdentityV1.from_bytes(raw)
    except ValueError as err:
        print(f"{DEVICE_ID_FILENAME}: {err}")
        return 2
    if not record.verify_crc(raw):
        print(f"{DEVICE_ID_FILENAME}: record CRC32 mismatch")
        return 2
    show("device", f"{record.device_model} {record.hw_revision}, serial {record.device_serial}")
    if record.timestamp:
        try:
            written = datetime.fromtimestamp(record.timestamp, timezone.utc)
        except (OverflowError, OSError, ValueError):
            # The signature does not cover the timestamp: any int64 fits.
            show("timestamp", f"{record.timestamp} (out of range)")
        else:
            show("timestamp", f"{written:%Y-%m-%d %H:%M:%S} UTC")

    algorithm = record.signature_algorithm
    if algorithm == SignatureAlgorithm.NONE:
        show("signature", "none (signature_version = 0)")
        return 2
    curve, key_file = ALGORITHMS[algorithm]
    key_path = args.key or Path(__file__).with_name(key_file)
    try:
        public_key = load_pem_public_key(key_path.read_bytes())
    except (OSError, ValueError, UnsupportedAlgorithm) as err:
        print(f"{key_path}: cannot load the public key: {err}")
        return 2
    if not isinstance(public_key, ec.EllipticCurvePublicKey) or public_key.curve.name != curve.name:
        print(f"{key_path}: not an EC {curve.name} public key")
        return 2

    payload = signed_payload(header)
    show("payload", payload.decode())

    # The record stores the raw r||s pair; cryptography expects DER.
    half = len(record.signature) // 2
    r = int.from_bytes(record.signature[:half], "big")
    s = int.from_bytes(record.signature[half:], "big")
    try:
        public_key.verify(encode_dss_signature(r, s), payload, ec.ECDSA(hashes.SHA256()))
    except InvalidSignature:
        show("signature", f"INVALID ({curve.name})")
        return 1
    show("signature", f"valid ({curve.name}, key {key_path.name})")
    return 0


if __name__ == "__main__":
    sys.exit(main())
