#!/usr/bin/env python3
"""Reference decryptor for BMENC v1 (docs/specs/bmenc/v1.md, formats section 11).

It exists so that the format is a format and not whatever the platform happens to write: a second implementation,
in another language, from the specification rather than from the C# source, checking the published vectors. If the
two disagree, one of them is wrong and the vectors say which. It is also the way into a `.bmenc` artifact that needs
nothing of the platform - not even bmctl - once the recovery key is in hand (formats section 24).

    python3 bmenc_ref.py [--vectors docs/specs/bmenc/vectors/v1/vectors.json] [--json]
    python3 bmenc_ref.py --identity recovery-key.txt artifact.tar.zst.bmenc > artifact.tar.zst
    python3 bmenc_ref.py --password-file pw.txt artifact.bmenc -o artifact

A file is decrypted one segment at a time, so memory stays at a segment whatever its size; every segment is
authenticated before its bytes are written, and a file cut short or with anything after its last segment fails.

Needs pyca/cryptography. Argon2id stanzas and vectors additionally need argon2-cffi; without it vectors are reported
as skipped rather than failed, because a reader that cannot open a password stanza is still a conforming reader of
the rest.
"""

from __future__ import annotations

import argparse
import hashlib
import hmac
import json
import pathlib
import sys

from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey, X25519PublicKey
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
from cryptography.hazmat.primitives.keywrap import aes_key_unwrap_with_padding, aes_key_wrap_with_padding

MAGIC = b"BMENC\x00"
VERSION = 1
AEAD_AES_GCM = 1
META_FLAG = 0x01
FIXED_HEADER = 48
MAC_LEN = 32
TAG_LEN = 16
KEY_LEN = 32
WRAPPED_LEN = 40
MIN_HEADER = 80
MAX_HEADER = 1024 * 1024
MAX_STANZAS = 32
MAX_META = 65536
MAX_ARGON_MEMORY_KIB = 4194304
MAX_ARGON_PASSES = 16
MAX_ARGON_LANES = 16
MAX_PBKDF2_ITERATIONS = 10_000_000

STANZA_KEK = 0x01
STANZA_ARGON2ID = 0x02
STANZA_PBKDF2 = 0x03
STANZA_X25519 = 0x04

BODY_LENGTHS = {STANZA_KEK: 60, STANZA_ARGON2ID: 65, STANZA_PBKDF2: 60, STANZA_X25519: 72}


class BmencError(Exception):
    """The file is not a BMENC v1 file, or does not verify."""


def hkdf(ikm: bytes, salt: bytes, info: bytes, length: int = KEY_LEN) -> bytes:
    return HKDF(algorithm=hashes.SHA256(), length=length, salt=salt, info=info).derive(ikm)


def schedule(file_key: bytes, stream_salt: bytes, aead_id: int, seg_log2: int) -> dict[str, bytes]:
    """The three subkeys of one file (section 11.3)."""
    return {
        "header": hkdf(file_key, stream_salt, b"BMENC/1/header"),
        "payload": hkdf(file_key, stream_salt, b"BMENC/1/payload" + bytes([aead_id, seg_log2])),
        "meta": hkdf(file_key, stream_salt, b"BMENC/1/meta"),
    }


def nonce(index: int, last: bool) -> bytes:
    """An 88-bit big-endian counter and the last-segment flag (section 11.4)."""
    return index.to_bytes(11, "big") + (b"\x01" if last else b"\x00")


def parse_header(data: bytes) -> dict:
    """Everything a reader checks before it has a key (section 11.6)."""
    if len(data) < 16:
        raise BmencError("the file is too short to be a BMENC file")
    if data[:6] != MAGIC:
        raise BmencError("not a BMENC file")
    if data[6] != VERSION:
        raise BmencError(f"unsupported BMENC version {data[6]} - upgrade bmctl")
    if data[7] != AEAD_AES_GCM:
        raise BmencError(f"unknown AEAD {data[7]}")
    seg_log2 = data[8]
    if not 16 <= seg_log2 <= 24:
        raise BmencError(f"segment size 2^{seg_log2} is outside the format")
    flags = data[9]
    if flags & ~META_FLAG:
        raise BmencError("undefined flag bit")
    count = int.from_bytes(data[10:12], "big")
    if not 1 <= count <= MAX_STANZAS:
        raise BmencError(f"a header carries 1 to 32 stanzas, not {count}")
    header_len = int.from_bytes(data[12:16], "big")
    if not MIN_HEADER <= header_len <= MAX_HEADER:
        raise BmencError(f"header length {header_len} is outside the format")
    if len(data) < header_len:
        raise BmencError("the file is shorter than its header claims")

    header = data[:header_len]
    end = header_len - MAC_LEN
    offset = FIXED_HEADER
    stanzas = []

    for _ in range(count):
        if offset + 3 > end:
            raise BmencError("a stanza runs past the end of the header")
        kind = header[offset]
        body_len = int.from_bytes(header[offset + 1:offset + 3], "big")
        if offset + 3 + body_len > end:
            raise BmencError("a stanza runs past the end of the header")
        if kind in BODY_LENGTHS and BODY_LENGTHS[kind] != body_len:
            raise BmencError(f"a stanza of type 0x{kind:02x} is {BODY_LENGTHS[kind]} bytes, not {body_len}")
        stanzas.append({"type": kind, "body": header[offset + 3:offset + 3 + body_len]})
        offset += 3 + body_len

    meta_offset = None
    meta_ct = None

    if flags & META_FLAG:
        if offset + 4 > end:
            raise BmencError("no room for the meta block the header claims")
        meta_offset = offset
        meta_len = int.from_bytes(header[offset:offset + 4], "big")
        if not TAG_LEN <= meta_len <= MAX_META or offset + 4 + meta_len > end:
            raise BmencError(f"the meta block claims {meta_len} bytes")
        meta_ct = header[offset + 4:offset + 4 + meta_len]
        offset += 4 + meta_len

    if offset != end:
        raise BmencError("the header has bytes in it that belong to nothing")

    return {
        "bytes": header,
        "length": header_len,
        "seg_log2": seg_log2,
        "flags": flags,
        "stream_salt": header[16:48],
        "stanzas": stanzas,
        "meta_offset": meta_offset,
        "meta_ct": meta_ct,
        "mac": header[end:],
    }


def unwrap(header: dict, keys: dict) -> tuple[bytes, int]:
    """The file key, from the first stanza a supplied key opens (section 11.6)."""
    for kind in (STANZA_KEK, STANZA_X25519, STANZA_ARGON2ID, STANZA_PBKDF2):
        for stanza in [s for s in header["stanzas"] if s["type"] == kind]:
            wrapping = wrapping_key(kind, stanza, keys)
            if wrapping is None:
                continue
            try:
                return aes_key_unwrap_with_padding(wrapping, stanza["body"][-WRAPPED_LEN:]), kind
            except Exception:  # noqa: BLE001 - a key that does not fit is not an error, it is the next stanza
                continue
    raise BmencError("none of the keys offered opens this file")


def wrapping_key(kind: int, stanza: dict, keys: dict) -> bytes | None:
    """What wraps the file key in this stanza, or None when no key for it was supplied (section 11.2)."""
    body = stanza["body"]

    if kind == STANZA_KEK:
        return keys.get("kek")

    if kind == STANZA_X25519:
        private = keys.get("recovery")
        if private is None:
            return None
        ephemeral_public = body[:32]
        shared = X25519PrivateKey.from_private_bytes(private).exchange(
            X25519PublicKey.from_public_bytes(ephemeral_public))
        if shared == bytes(32):
            raise BmencError("the recovery stanza agreed on nothing")
        recipient_public = X25519PrivateKey.from_private_bytes(private).public_key().public_bytes_raw()
        return hkdf(shared, ephemeral_public + recipient_public, b"BMENC/1/X25519")

    if kind == STANZA_ARGON2ID:
        password = keys.get("password")
        if password is None:
            return None
        memory = int.from_bytes(body[16:20], "big")
        passes = int.from_bytes(body[20:24], "big")
        lanes = body[24]
        if not 0 < memory <= MAX_ARGON_MEMORY_KIB or not 0 < passes <= MAX_ARGON_PASSES \
                or not 0 < lanes <= MAX_ARGON_LANES:
            raise BmencError("the password stanza asks for more work than this reader will do")
        try:
            from argon2.low_level import Type, hash_secret_raw
        except ImportError as missing:  # pragma: no cover - depends on the machine
            raise Skipped("argon2-cffi is not installed") from missing
        return hash_secret_raw(
            secret=password.encode("utf-8"), salt=body[:16], time_cost=passes, memory_cost=memory,
            parallelism=lanes, hash_len=KEY_LEN, type=Type.ID, version=0x13)

    if kind == STANZA_PBKDF2:
        password = keys.get("password")
        if password is None:
            return None
        iterations = int.from_bytes(body[16:20], "big")
        if not 0 < iterations <= MAX_PBKDF2_ITERATIONS:
            raise BmencError("the password stanza asks for more work than this reader will do")
        return PBKDF2HMAC(
            algorithm=hashes.SHA256(), length=KEY_LEN, salt=body[:16], iterations=iterations,
        ).derive(password.encode("utf-8"))

    return None


class Skipped(Exception):
    """This machine cannot check this vector; it is not a failure of the format."""


def decrypt(data: bytes, keys: dict) -> tuple[bytes, dict | None]:
    """The plaintext of a whole file, and what it says about itself (section 11.6)."""
    header = parse_header(data)
    file_key, _ = unwrap(header, keys)
    subkeys = schedule(file_key, header["stream_salt"], AEAD_AES_GCM, header["seg_log2"])

    expected = hmac.new(subkeys["header"], header["bytes"][:-MAC_LEN], hashlib.sha256).digest()
    if not hmac.compare_digest(expected, header["mac"]):
        raise BmencError("the header does not verify; it has been tampered with")

    meta = None
    if header["meta_ct"] is not None:
        aad = header["bytes"][:header["meta_offset"]]
        meta = json.loads(AESGCM(subkeys["meta"]).decrypt(bytes(12), header["meta_ct"], aad))

    stride = (1 << header["seg_log2"]) + TAG_LEN
    payload = data[header["length"]:]
    if not payload:
        raise BmencError("the file has a header and no payload")

    gcm = AESGCM(subkeys["payload"])
    plaintext = bytearray()
    index = 0
    offset = 0

    while offset < len(payload):
        chunk = payload[offset:offset + stride]
        last = offset + len(chunk) >= len(payload)
        if len(chunk) < TAG_LEN:
            raise BmencError("the payload ends inside a segment; it is truncated")
        if len(chunk) == TAG_LEN and index > 0:
            raise BmencError("a segment with no bytes in it is only ever the whole of an empty file")
        try:
            plaintext += gcm.decrypt(nonce(index, last), chunk, None)
        except Exception as failure:  # noqa: BLE001 - every AEAD failure is the same failure
            raise BmencError(f"segment {index} does not verify") from failure
        offset += len(chunk)
        index += 1

    return bytes(plaintext), meta


def build(vector: dict) -> bytes:
    """The whole file of a KEK-only vector, built from its inputs: the byte-for-byte agreement check."""
    file_key = bytes.fromhex(vector["fileKey"])
    stream_salt = bytes.fromhex(vector["streamSalt"])
    seg_log2 = vector["segLog2"]
    subkeys = schedule(file_key, stream_salt, AEAD_AES_GCM, seg_log2)

    kek = vector["kek"]
    body = uuid_bytes(kek["id"]) + kek["version"].to_bytes(4, "big")
    body += aes_key_wrap_with_padding(bytes.fromhex(kek["key"]), file_key)
    stanza = bytes([STANZA_KEK]) + len(body).to_bytes(2, "big") + body

    header_len = FIXED_HEADER + len(stanza) + MAC_LEN
    header = bytearray(MAGIC + bytes([VERSION, AEAD_AES_GCM, seg_log2, 0]))
    header += (1).to_bytes(2, "big") + header_len.to_bytes(4, "big") + stream_salt + stanza
    header += hmac.new(subkeys["header"], bytes(header), hashlib.sha256).digest()

    gcm = AESGCM(subkeys["payload"])
    plaintext = pattern(vector["plaintextLength"])
    size = 1 << seg_log2
    segments = [plaintext[at:at + size] for at in range(0, len(plaintext), size)] or [b""]

    out = bytearray(header)
    for index, segment in enumerate(segments):
        out += gcm.encrypt(nonce(index, index == len(segments) - 1), segment, None)

    return bytes(out)


def uuid_bytes(value: str) -> bytes:
    """A UUID in RFC 9562 byte order, which is what the KEK stanza carries."""
    return bytes.fromhex(value.replace("-", ""))


def pattern(length: int) -> bytes:
    """The plaintext of every vector: (i * 131) % 251."""
    return bytes((index * 131) % 251 for index in range(length))


def mutate(data: bytes, negative: dict, header: dict) -> bytes:
    """The change a negative vector names (section 11.11)."""
    kind = negative["kind"]
    offset = negative["offset"]
    if offset < 0:
        offset += len(data)

    out = bytearray(data)
    if kind == "flip":
        out[offset] ^= 0x01
    elif kind == "set":
        out[offset] = negative["value"]
    elif kind == "truncate":
        del out[len(out) - negative["value"]:]
    elif kind == "append":
        out += bytes([0x7E]) * negative["value"]
    elif kind == "flipMeta":
        out[header["meta_offset"] + 4] ^= 0x01
    elif kind == "flipHeaderMac":
        out[header["length"] - 1] ^= 0x01
    else:
        raise ValueError(f"unknown mutation {kind}")

    return bytes(out)


def keys_of(vector: dict) -> list[dict]:
    """Every key the vector says must open it."""
    keys = []
    if vector.get("kek"):
        keys.append({"kek": bytes.fromhex(vector["kek"]["key"])})
    if vector.get("recovery"):
        keys.append({"recovery": bytes.fromhex(vector["recovery"]["privateKey"])})
    if vector.get("password"):
        keys.append({"password": vector["password"]["value"]})
    return keys


def check(path: pathlib.Path) -> dict:
    """Every vector in the file, checked; the report says what happened to each kind."""
    vectors = json.loads(path.read_text(encoding="utf-8"))
    report = {"decrypted": 0, "rebuilt": 0, "refused": 0, "skipped": [], "failed": []}

    for vector in vectors["positive"]:
        data = bytes.fromhex(vector["fileHex"])
        expected = pattern(vector["plaintextLength"])

        if hashlib.sha256(data).hexdigest() != vector["fileSha256"]:
            report["failed"].append(f"{vector['name']}: the published digest is not the published bytes")
            continue

        for keys in keys_of(vector):
            try:
                plaintext, _ = decrypt(data, keys)
            except Skipped as skipped:
                report["skipped"].append(f"{vector['name']}: {skipped}")
                continue
            except BmencError as failure:
                report["failed"].append(f"{vector['name']}: {failure}")
                continue

            if plaintext != expected:
                report["failed"].append(f"{vector['name']}: decrypted to something else")
            else:
                report["decrypted"] += 1

        if vector.get("kek") and not vector.get("recovery") and not vector.get("password") \
                and not vector.get("meta") and not vector.get("built"):
            if build(vector) == data:
                report["rebuilt"] += 1
            else:
                report["failed"].append(f"{vector['name']}: this reference builds different bytes")

    for vector in vectors["boundary"]:
        built = build(vector)
        if hashlib.sha256(built).hexdigest() == vector["fileSha256"]:
            report["rebuilt"] += 1
        else:
            report["failed"].append(f"{vector['name']}: this reference builds different bytes")

    positive = {vector["name"]: vector for vector in vectors["positive"]}

    for negative in vectors["negative"]:
        vector = positive[negative["of"]]
        data = bytes.fromhex(vector["fileHex"])
        broken = mutate(data, negative, parse_header(data))

        try:
            decrypt(broken, keys_of(vector)[0])
        except Skipped as skipped:
            report["skipped"].append(f"{negative['name']}: {skipped}")
        except BmencError:
            report["refused"] += 1
        else:
            report["failed"].append(f"{negative['name']}: this reference accepted a file it must refuse")

    return report


BECH32_CHARSET = "qpzry9x8gf2tvdw0s3jn54khce6mua7l"
IDENTITY_PREFIX = "age-secret-key-"


def bech32_polymod(values: list[int]) -> int:
    """The BIP-173 checksum polynomial, which age uses for its keys."""
    generator = (0x3B6A57B2, 0x26508E6D, 0x1EA119FA, 0x3D4233DD, 0x2A1462B3)
    checksum = 1
    for value in values:
        top = checksum >> 25
        checksum = ((checksum & 0x1FFFFFF) << 5) ^ value
        for bit in range(5):
            if (top >> bit) & 1:
                checksum ^= generator[bit]
    return checksum


def identity_key(text: str) -> bytes | None:
    """The X25519 private key of an `AGE-SECRET-KEY-1...` line, or None when the line is not one."""
    line = text.strip()
    if not line.upper().startswith(IDENTITY_PREFIX.upper() + "1") or (line != line.upper() and line != line.lower()):
        return None

    lower = line.lower()
    separator = lower.rfind("1")
    prefix, data = lower[:separator], lower[separator + 1:]
    if prefix != IDENTITY_PREFIX or len(data) < 7 or any(char not in BECH32_CHARSET for char in data):
        return None

    values = [BECH32_CHARSET.index(char) for char in data]
    expanded = [ord(char) >> 5 for char in prefix] + [0] + [ord(char) & 31 for char in prefix]
    if bech32_polymod(expanded + values) != 1:
        return None

    accumulator, bits, key = 0, 0, bytearray()
    for value in values[:-6]:
        accumulator = (accumulator << 5) | value
        bits += 5
        if bits >= 8:
            bits -= 8
            key.append((accumulator >> bits) & 0xFF)
    if bits >= 5 or accumulator & ((1 << bits) - 1):
        return None

    return bytes(key) if len(key) == KEY_LEN else None


def read_identity(path: pathlib.Path) -> bytes:
    """The recovery private key in an identity file: the first `AGE-SECRET-KEY-1...` line; comments are skipped."""
    for line in path.read_text(encoding="utf-8").splitlines():
        key = identity_key(line)
        if key is not None:
            return key
    raise BmencError(f"{path} holds no AGE-SECRET-KEY-1 line")


def read_exactly(source, count: int) -> bytes:
    """Up to `count` bytes; fewer only at the end of the input, as a pipe may hand them over in pieces."""
    chunks = []
    wanted = count
    while wanted > 0:
        chunk = source.read(wanted)
        if not chunk:
            break
        chunks.append(chunk)
        wanted -= len(chunk)
    return b"".join(chunks)


def decrypt_stream(source, sink, keys: dict) -> dict | None:
    """
    Decrypts a whole file from `source` into `sink`, a segment at a time (section 11.6): the header is checked before
    any key is tried, each segment is authenticated before a byte of it is written, and the segment after it is read
    first, because only the last segment says it is last. Returns what the file says about itself.
    """
    fixed = read_exactly(source, 16)
    if len(fixed) < 16:
        raise BmencError("the file is too short to be a BMENC file")
    header_len = int.from_bytes(fixed[12:16], "big")
    if not MIN_HEADER <= header_len <= MAX_HEADER:
        parse_header(fixed)  # names what is wrong with the fixed part, the length included
        raise BmencError(f"header length {header_len} is outside the format")

    header_bytes = fixed + read_exactly(source, header_len - 16)
    header = parse_header(header_bytes)
    file_key, _ = unwrap(header, keys)
    subkeys = schedule(file_key, header["stream_salt"], AEAD_AES_GCM, header["seg_log2"])

    expected = hmac.new(subkeys["header"], header["bytes"][:-MAC_LEN], hashlib.sha256).digest()
    if not hmac.compare_digest(expected, header["mac"]):
        raise BmencError("the header does not verify; it has been tampered with")

    meta = None
    if header["meta_ct"] is not None:
        aad = header["bytes"][:header["meta_offset"]]
        meta = json.loads(AESGCM(subkeys["meta"]).decrypt(bytes(12), header["meta_ct"], aad))

    stride = (1 << header["seg_log2"]) + TAG_LEN
    gcm = AESGCM(subkeys["payload"])
    segment = read_exactly(source, stride)
    if not segment:
        raise BmencError("the file has a header and no payload")

    index = 0
    while True:
        following = read_exactly(source, stride)
        last = not following
        if len(segment) < TAG_LEN:
            raise BmencError("the payload ends inside a segment; it is truncated")
        if len(segment) == TAG_LEN and index > 0:
            raise BmencError("a segment with no bytes in it is only ever the whole of an empty file")
        try:
            plaintext = gcm.decrypt(nonce(index, last), segment, None)
        except Exception as failure:  # noqa: BLE001 - every AEAD failure is the same failure
            raise BmencError(f"segment {index} does not verify") from failure
        sink.write(plaintext)
        if last:
            return meta
        segment = following
        index += 1


def decrypt_file(arguments: argparse.Namespace) -> int:
    """`--identity` or `--password-file` with a file: the plaintext to `-o`, or to standard output."""
    keys = {}
    if arguments.identity:
        keys["recovery"] = read_identity(arguments.identity)
    if arguments.password_file:
        lines = arguments.password_file.read_text(encoding="utf-8").splitlines()
        keys["password"] = lines[0] if lines else ""

    source = sys.stdin.buffer if str(arguments.input) == "-" else open(arguments.input, "rb")  # noqa: SIM115
    partial = None
    try:
        if arguments.output and str(arguments.output) != "-":
            # Written under another name and moved into place once every segment verified.
            partial = arguments.output.with_name(arguments.output.name + ".partial")
            with open(partial, "wb") as sink:
                decrypt_stream(source, sink, keys)
            partial.replace(arguments.output)
            partial = None
        else:
            decrypt_stream(source, sys.stdout.buffer, keys)
            sys.stdout.buffer.flush()
    except (BmencError, Skipped) as failure:
        print(f"bmenc_ref: {failure}", file=sys.stderr)
        return 1
    finally:
        if source is not sys.stdin.buffer:
            source.close()
        if partial is not None and partial.exists():
            partial.unlink()

    return 0


def main() -> int:
    parser = argparse.ArgumentParser(
        description="Decrypt a BMENC v1 file with the recovery key or a password, or check the conformance vectors.")
    parser.add_argument("input", nargs="?", type=pathlib.Path, help="a .bmenc file to decrypt, or - for standard input")
    parser.add_argument("--identity", type=pathlib.Path, help="the recovery key file (AGE-SECRET-KEY-1...): opens stanza 0x04")
    parser.add_argument("--password-file", type=pathlib.Path, help="a password on the file's first line: opens stanzas 0x02 and 0x03")
    parser.add_argument("-o", "--output", type=pathlib.Path, help="where the plaintext goes; standard output by default")
    parser.add_argument(
        "--vectors",
        type=pathlib.Path,
        default=pathlib.Path(__file__).resolve().parent.parent / "vectors" / "v1" / "vectors.json")
    parser.add_argument("--json", action="store_true", help="print the vector report as JSON")
    arguments = parser.parse_args()

    if arguments.input is not None:
        if not arguments.identity and not arguments.password_file:
            parser.error("decrypting a file needs --identity or --password-file")
        return decrypt_file(arguments)

    report = check(arguments.vectors)

    if arguments.json:
        print(json.dumps(report, indent=2))
    else:
        print(f"decrypted {report['decrypted']}, rebuilt {report['rebuilt']}, refused {report['refused']}")
        for line in report["skipped"]:
            print(f"skipped: {line}")
        for line in report["failed"]:
            print(f"FAILED: {line}")

    return 1 if report["failed"] else 0


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