#!/usr/bin/env python3

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

asnoid_sha256 = "3031300d060960864801650304020105000420"

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 SHA256 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_sha256) + d)
    if p.returncode != 0:
        raise Exception("signing failed")
    return output

if len(sys.argv) < 3:
    print("usage: %s <key> <filename>" % sys.argv[0])
    sys.exit(1)

key_file = sys.argv[1]
fn = sys.argv[2]

SUPERBLOCK_SIZE = 96

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

    if len(superblock) < SUPERBLOCK_SIZE:
        sys.stderr.write("Unable to read superblock. Perhaps image is not a squashfs?\n")
        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:
        sys.stderr.write("Image does not appear to be a squashfs. Magic is 0x%x.\n" % magic)
        sys.exit(1)

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

    if os.path.getsize(fn) < bytes_used:
        sys.stderr.write("Image has been truncated\n")
        sys.exit(1)

    # mksquashfs pads the image to 4K so the size might differ by that much
    if os.path.getsize(fn) > bytes_used + 4095:
        sys.stderr.write("Image is already signed or has extra junk on the end.  Not signing it again.\n")
        sys.exit(1)

    sha = hashlib.sha256()
    sha.update(superblock)

    # Now read the remaining used bytes and hash them.
    remainder = f.read(bytes_used - SUPERBLOCK_SIZE);

    if len(remainder) + SUPERBLOCK_SIZE < bytes_used:
        sys.stderr.write("Unable to read entire body of image\n")
        sys.exit(1)

    sha.update(remainder)
    digest = sha.digest()

    signature = sign_hash(digest, key_file)

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

    print("Signature length: %u" % len(signature))

    # When using stdio there must be an intervening seek
    f.seek(0, 1)

    f.write(b"BSSG")
    f.write(signature_length)
    f.write(signature)

    # We may have extended the file by writing the signature which
    # means that it will no longer be a multiple of 4K. Let's pad it
    # so that it is.
    while (f.tell() % 4096):
        f.write(b'\xFF')

    f.close()
    print("Signature: %s" % signature.hex())
    print("Filename: %s" % fn)
