"""Verify a release signature against the pinned public key.

WHY THIS IS HAND-WRITTEN, WHICH NORMALLY IT SHOULD NOT BE

"Don't roll your own crypto" is right, and this is one of the narrow cases where
the reasoning behind the rule does not apply.

The dangerous parts of RSA are on the *signing* side — timing and side channels
around a secret exponent. There is no secret here. This module holds a public
key and answers one question: do these bytes carry a signature made by the
offline key? Everything it touches is already public.

The one classic failure in a verifier is Bleichenbacher's forgery, which works
when an implementation *parses* the padded block leniently and ignores trailing
data. This one never parses. It reconstructs the entire expected block from the
digest and compares all of it, so there is nowhere for extra bytes to hide.

What it buys: the installer verifies a package on a bare Python, before any
dependency has been installed, on a machine that may have just been rebuilt.
Making the most security-critical step in the deployment depend on a compiled
wheel being present is a worse trade than 60 lines of arithmetic.

What makes it trustworthy is not this docstring. `tests/test_release_signing.py`
verifies signatures produced by PHP's OpenSSL against the real 4096-bit release
key — the same cross-language discipline the HMAC fixtures use — and asserts
that tampered versions, swapped signatures and corrupted padding are all
refused.
"""

from __future__ import annotations

import base64
import binascii
import hashlib
import hmac

#: The ASN.1 DigestInfo prefix for SHA-256, from RFC 8017 §9.2 notes.
#:
#: A constant rather than something assembled at runtime: it is a fixed byte
#: string for a fixed algorithm, and building it dynamically would mean parsing
#: attacker-supplied ASN.1 to decide what to compare against — which is the
#: mistake this whole module is written to avoid.
SHA256_DIGEST_INFO = bytes.fromhex("3031300d060960864801650304020105000420")


class SignatureError(Exception):
    """The signature could not be checked — malformed key, or unusable input.

    Distinct from a signature that is simply wrong, which is a `False` return.
    An operator needs to tell "this package was not signed by us" from "this
    installation cannot verify anything", because the second is a broken
    deployment and the first is an attack or a mistake.
    """


class PublicKey:
    """An RSA public key, parsed once."""

    __slots__ = ("modulus", "exponent", "size")

    def __init__(self, modulus: int, exponent: int) -> None:
        if modulus <= 0 or exponent <= 0:
            raise SignatureError("The public key is not usable.")

        self.modulus = modulus
        self.exponent = exponent
        #: Length of the modulus in bytes — `k` in RFC 8017.
        self.size = (modulus.bit_length() + 7) // 8

        if self.size < 256:
            # 2048 bits. Below that a signature is not worth checking, and the
            # release key is 4096. This refuses rather than warns because a
            # weak key that verifies is worse than one that does not load.
            raise SignatureError(f"The public key is too small ({self.size * 8} bits).")

    def __repr__(self) -> str:  # pragma: no cover - diagnostics only
        return f"PublicKey(bits={self.size * 8}, e={self.exponent})"


def verify(message: str, signature: str, public_key: str) -> bool:
    """Was `message` signed by the private key matching `public_key`?

    `signature` is base64, as the License Server stores and serves it.
    `message` is the canonical manifest string — see manifest.py.

    Returns False for a bad signature. Raises SignatureError when verification
    could not be performed at all.
    """
    key = load_public_key(public_key)

    try:
        raw = base64.b64decode(signature.strip(), validate=True)
    except (binascii.Error, ValueError):
        # Not an exception: a corrupted signature field is exactly the case
        # this exists to reject, and it must read as "refused", not "broken".
        return False

    if len(raw) != key.size:
        # RFC 8017 §8.2.2 step 1. A signature is always exactly the modulus
        # length; anything else cannot be right, and short values are how
        # padding forgeries are smuggled in.
        return False

    encoded = int.from_bytes(raw, "big")

    if encoded >= key.modulus:
        return False

    recovered = pow(encoded, key.exponent, key.modulus).to_bytes(key.size, "big")
    expected = _encoded_message(message, key.size)

    # Whole-block comparison. Nothing here parses the recovered block, looks for
    # a separator, or reads a length — the padding, the DigestInfo and the hash
    # are all rebuilt and compared together, which is what closes the forgery.
    return hmac.compare_digest(recovered, expected)


def load_public_key(pem: str) -> PublicKey:
    """Parse a PEM public key.

    Accepts both `BEGIN PUBLIC KEY` (SubjectPublicKeyInfo, what OpenSSL and PHP
    produce by default, and what the release key is) and `BEGIN RSA PUBLIC KEY`
    (a bare PKCS#1 RSAPublicKey), because a key generated by an older toolchain
    should not be a deployment failure nobody can diagnose.
    """
    der = _pem_body(pem)
    tag, value, _ = _read_tlv(der, 0)

    if tag != 0x30:
        raise SignatureError("The public key is not a DER sequence.")

    # PKCS#1: SEQUENCE { INTEGER modulus, INTEGER exponent }.
    # SPKI:   SEQUENCE { SEQUENCE algorithm, BIT STRING key }.
    inner_tag, _, offset = _read_tlv(value, 0)

    if inner_tag == 0x02:
        return _rsa_public_key(value)

    if inner_tag != 0x30:
        raise SignatureError("Unrecognised public key structure.")

    bit_tag, bit_value, _ = _read_tlv(value, offset)

    if bit_tag != 0x03 or not bit_value or bit_value[0] != 0x00:
        # The leading byte of a BIT STRING counts unused trailing bits. A key
        # is whole bytes, so anything but zero means this is not what we think.
        raise SignatureError("Unrecognised public key structure.")

    key_tag, key_value, _ = _read_tlv(bit_value[1:], 0)

    if key_tag != 0x30:
        raise SignatureError("Unrecognised public key structure.")

    return _rsa_public_key(key_value)


# ── Internals ─────────────────────────────────────────────────────────────


def _encoded_message(message: str, size: int) -> bytes:
    """Build EMSA-PKCS1-v1_5, RFC 8017 §9.2.

        0x00 || 0x01 || 0xFF... || 0x00 || DigestInfo || hash
    """
    digest = hashlib.sha256(message.encode("utf-8")).digest()
    tail = SHA256_DIGEST_INFO + digest

    # §9.2 step 3: at least 8 bytes of padding, or the key is too small for the
    # algorithm. Unreachable at 2048 bits and above, which the key size check
    # already enforces — kept because the constant it depends on is elsewhere.
    if size < len(tail) + 11:
        raise SignatureError("The public key is too small for SHA-256 signatures.")

    return b"\x00\x01" + b"\xff" * (size - len(tail) - 3) + b"\x00" + tail


def _pem_body(pem: str) -> bytes:
    lines = [line.strip() for line in pem.strip().splitlines()]
    body = "".join(line for line in lines if line and not line.startswith("-----"))

    if not body:
        raise SignatureError("The public key is empty.")

    try:
        return base64.b64decode(body, validate=True)
    except (binascii.Error, ValueError) as exc:
        raise SignatureError("The public key is not valid base64.") from exc


def _rsa_public_key(der: bytes) -> PublicKey:
    modulus_tag, modulus, offset = _read_tlv(der, 0)
    exponent_tag, exponent, _ = _read_tlv(der, offset)

    if modulus_tag != 0x02 or exponent_tag != 0x02:
        raise SignatureError("The public key does not hold two integers.")

    return PublicKey(
        int.from_bytes(modulus, "big"),
        int.from_bytes(exponent, "big"),
    )


def _read_tlv(data: bytes, offset: int) -> tuple[int, bytes, int]:
    """One DER tag-length-value. Returns (tag, value, offset after it).

    Only what a public key needs: definite lengths, no nesting logic, no tag
    classes. Anything unexpected raises rather than being interpreted — this
    reads a key we ship ourselves, and leniency here would be leniency about
    what is allowed to verify a release.
    """
    if offset + 2 > len(data):
        raise SignatureError("The public key is truncated.")

    tag = data[offset]
    length = data[offset + 1]
    offset += 2

    if length & 0x80:
        count = length & 0x7F

        if count == 0 or count > 4 or offset + count > len(data):
            raise SignatureError("The public key has an unreadable length.")

        length = int.from_bytes(data[offset : offset + count], "big")
        offset += count

    if offset + length > len(data):
        raise SignatureError("The public key is truncated.")

    return tag, data[offset : offset + length], offset + length
