#!/usr/bin/env python3
# -*- coding: utf-8 -*-
#
# Written for the OpenSSL appendix in the CrypTool book.
#
# (c) 2020, 2022 Bernd Busse / BE  [open source Apache 2 licence]
# Version: 1.0.0 (2020-06-03)
#          1.0.1 (2022-03-21)
"""Convert number to/from file for de-/encryption with `openssl rsautl`

This file does only the format conversion. It doesn't need any pem files.

Workflow (together with the calling shell script):

Convert plain integer into binary format for direct use with `openssl rsautl`:
1. Create *plaintext* input (filled with null-bytes as needed) with:
  `python int2rsa.py -enc NUMBER KEYSIZE message.bin`
2. Encrypt (raw RSA, no padding) with:
  `openssl rsautl -encrypt -raw -inkey pubkey.pem -pubin -in message.bin -out message.bin.enc`

Convert binary output from `openssl rsautl` to plain integer:
1. Decrypt (raw RSA, no padding) with:
  `openssl rsautl -decrypt raw -inkey key.pem -in message.bin.enc -out message.bin.dec`
2. Convert *decrypted* output with:
  `python int2rsa.py -dec message.bin.dec`
"""  # noqa: E501

import os.path
import sys


class ConversionError(RuntimeError):
    """Generic Error while converting."""
    pass


def print_err(msg, *args, **kwargs):
    """Print formatted `msg` to `sys.stderr`."""
    print(msg.format(*args, **kwargs), file=sys.stderr)


def print_usage(file=sys.stdout):
    """Print usage message."""
    print(f"usage: {sys.argv[0]} {{ -enc ENC_OPTIONS | -dec DEC_OPTIONS }}\n"
          f" ENC_OPTIONS := NUMBER KEYSIZE OUTPUT_FILE\n"
          f" DEC_OPTIONS := INPUT_FILE\n",
          file=file)


def convert_int2bin(number, keysize, outfile):
    """Write `number` to `out` in OpenSSL binary format."""
    if keysize % 8 != 0:
        raise ConversionError("KEYSIZE must be a multiple of 8")

    try:
        num_bytes = keysize // 8
        bin_data = number.to_bytes(num_bytes, "big")
    except OverflowError as err:
        raise ConversionError(
            f"{number} to big to fit into {num_bytes} bytes"
        ) from err

    if os.path.exists(outfile):
        print(f"Warn: '{outfile}' already exists.", end=" ")
        prompt = input("Overwrite? [y|N] ")
        if prompt.lower() not in ("y", "yes"):
            raise ConversionError(f"Cannot write binary output: "
                                  f"File '{outfile}' already exists")
    try:
        with open(outfile, "wb") as binary:
            binary.write(bin_data)
    except OSError as err:
        raise ConversionError(f"Cannot write binary output: {err!s}") from err


def convert_bin2int(infile):
    """Convert `infile` from OpenSSL binary format to integer."""
    try:
        with open(infile, "rb") as binary:
            bin_data = binary.read()
            number = int.from_bytes(bin_data, "big")
            print(number)
    except OSError as err:
        raise ConversionError(f"Cannot open binary input: {err!s}") from err


def main():
    """Handle interactive invocation."""
    if len(sys.argv) < 2:
        print_usage(sys.stderr)
        return False

    mode = sys.argv[1]
    if mode == "-enc":
        if len(sys.argv) != 5:
            print_err("Error: Invalid parameters for encryption conversion")
            print_usage(sys.stderr)
            return False

        try:
            number = int(sys.argv[2])
        except ValueError as err:
            print_err(f"Error: invalid value for parameter NUMBER: {err!s}")
            return False
        try:
            keysize = int(sys.argv[3])
        except ValueError as err:
            print_err(f"Error: invalid value for parameter KEYSIZE: {err!s}")
            return False

        convert_int2bin(number, keysize, sys.argv[4])

    elif mode == "-dec":
        if len(sys.argv) != 3:
            print_err("Error: Invalid parameters for decryption conversion")
            print_usage(sys.stderr)
            return False

        convert_bin2int(sys.argv[2])

    elif mode in ("-h", "-help", "--help"):
        print_usage()
        return True

    else:
        print_err(f"Error: Unsupported mode '{mode}'.")
        print_usage(sys.stderr)
        return False

    return True


if __name__ == '__main__':
    try:
        if not main():
            exit(1)
    except ConversionError as err:
        print_err(f"Error: Conversion failed: {err!s}")
        exit(1)
