#!/usr/bin/env python3

# Generates a Brightsign Verity Signed (bsvs) file from a squashfs
# image. A custom BS header is appended to the squashfs image which
# contains a signature constructed from the verity root hash of the
# squashfs image and the hash device which is the dm-verity header
# and root hash tree, see [1]. The remainder of the header is the
# actual root hash, root hash device offset and actual root hash
# device data.

# [1] https://gitlab.com/cryptsetup/cryptsetup/-/wikis/DMVerity

import sys, os, struct, hashlib, subprocess, shlex

asnoid_sha512 = "3051300d060960864801650304020305000440"

def sign_hash(d, key_file):
    cmd = "openssl rsautl -sign -pkcs -inkey " + key_file
    p = subprocess.Popen(shlex.split(cmd), stdin=subprocess.PIPE, stdout=subprocess.PIPE)

    # Apply PKCS#1 standard padding for SHA512 which is used to automatically determine the
    # hash used by libgcrypt. The remaining padding will be applied by rsautl.
    (output, err) = p.communicate(bytes.fromhex(asnoid_sha512) + d)
    if p.returncode != 0:
        raise Exception("signing failed")
    return output

if len(sys.argv) < 4:
    print(("usage: %s <key> <output_filename> <hash_device_filename> <root_hash>" % sys.argv[0]))
    sys.exit(1)

key_file = sys.argv[1]
fn = sys.argv[2]
hash_device_fn = sys.argv[3]
root_hash = bytes.fromhex(sys.argv[4])

SUPERBLOCK_SIZE = 96

with open(fn, 'rb+') as f:
    superblock = f.read(SUPERBLOCK_SIZE)

    if len(superblock) < SUPERBLOCK_SIZE:
        print("Unable to read superblock. Perhaps image is not a squashfs?")
        sys.exit(1)

    # From struct squashfs_super_block in linux/fs/squashfs/squashfs_fs.h
    (magic, inodes, mkfs_time, block_size, fragments, compression, block_log, flags, no_ids, s_major, s_minor, root_inode, bytes_used) = struct.unpack("<IIIIIHHHHHHQQ", superblock[0:48])

    if magic != 0x73717368:
        print("Image does not appear to be a squashfs. Magic is 0x%x." % magic)
        sys.exit(1)

    print("magic = %u" % magic)
    print("bytes_used = %Lu" % bytes_used)

    # mksquashfs pads to 4k alignment so add the padding on if we're not on
    # a 4k boundary, this is the location to  start writing our signature without
    # invalidating the dm-verity calculation which includes the padding up to 4k.
    bytes_used_padded = (bytes_used + (4096 - 1)) & ~(4096-1)

    if os.path.getsize(fn) < bytes_used_padded:
        print("Image has been truncated")
        sys.exit(1)

    with open(hash_device_fn, 'rb') as hf:
        hash_device = hf.read()

    # signature
    # -------------------------------------------------------------------
    # | root_hash_length | root_hash | hash_device_length | hash_device |
    # -------------------------------------------------------------------

    sha = hashlib.sha512()

    root_hash_length = struct.pack("<I", len(root_hash))
    hash_device_length = struct.pack("<I", len(hash_device))

    sha.update(root_hash_length)
    sha.update(root_hash)
    sha.update(hash_device_length)
    sha.update(hash_device)

    digest = sha.digest()

    signature = sign_hash(digest, key_file)

    signature_length = struct.pack("<I", len(signature))

    print("Signature length: %u" % len(signature))
    print("Root hash length: %u" % len(root_hash))
    print("Hash device length: %u" % len(hash_device))

    f.seek(bytes_used_padded)

    # header
    # -------------------------------------------------------------------------------------------------------------------------
    # |    -     | 4 bytes | uint32  |  -  | uint32        |     -     | uint32           | uint32     |   -    |      -      |
    # -------------------------------------------------------------------------------------------------------------------------
    # | squashfs |  "BSVS" | sig_len | sig | root_hash_len | root_hash | hash_device_size | pad_amount | pad 4k | hash_device |
    # -------------------------------------------------------------------------------------------------------------------------

    f.write(b"BSVS")

    f.write(signature_length)
    f.write(signature)

    f.write(root_hash_length)
    f.write(root_hash)

    f.write(hash_device_length)

    uint32_size = struct.calcsize("<I")

    # Check if we roll over the 4k boundary with pad_amount, if so
    # calculate the remainder based on the 4k boundary roll over
    pad_amount = 4096 - (f.tell() + uint32_size) % 4096
    if (pad_amount == 4096):
        pad_amount = 0

    # 4k pad amount
    f.write(struct.pack("<I", pad_amount))

    # Push to a 4k boundary as dm-verity metadata requires 4k alignment
    while (f.tell() % 4096):
        f.write(b'\xFF')

    hash_device_start = f.tell()

    # hash device (metadata)
    f.write(hash_device)

    f.close()

    print("Signature: %s" % signature.hex())
    print("Root Hash: %s" % root_hash.hex())
    print("Hash Device Start: %d" % hash_device_start)
