"""An RSA keypair for tests, generated on the spot.

WHY NOT JUST COMMIT A TEST KEY

Because this project's entire release story is "the private key has never been
in a repository", and a file called `test-key-private.pem` sitting in `tests/`
undermines that at a glance — for a reader, for a secret scanner, and for
whoever later needs a key in a hurry and finds one already checked in.

So the tests make their own, use it, and throw it away. It costs a couple of
seconds once per session (the fixture is session-scoped) and nothing is left
behind.

There is a second benefit. The DER *encoder* here is written independently of
the DER *parser* in app/release/signing.py, so a key produced by this module and
read by that one exercises both against each other rather than against a shared
assumption. The OpenSSL fixture proves the parser handles real keys; this proves
it handles a key it did not produce.

Nothing here is used outside tests, and none of it is written to be
constant-time or otherwise safe for real keys.
"""

from __future__ import annotations

import base64
import hashlib
import math
import random

#: DER for AlgorithmIdentifier { rsaEncryption, NULL } — a fixed byte string.
RSA_ALGORITHM_ID = bytes.fromhex("300d06092a864886f70d0101010500")

SHA256_DIGEST_INFO = bytes.fromhex("3031300d060960864801650304020105000420")

#: Trial division before Miller-Rabin. Removes about four candidates in five for
#: the cost of a few remainders, which is what keeps generation to seconds
#: rather than minutes in pure Python.
SMALL_PRIMES = [
    2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37, 41, 43, 47, 53, 59, 61, 67, 71,
    73, 79, 83, 89, 97, 101, 103, 107, 109, 113, 127, 131, 137, 139, 149, 151,
    157, 163, 167, 173, 179, 181, 191, 193, 197, 199, 211, 223, 227, 229, 233,
    239, 241, 251, 257, 263, 269, 271, 277, 281, 283, 293,
]


class TestKey:
    """A keypair, able to sign the way the offline key does."""

    def __init__(self, modulus: int, exponent: int, private: int) -> None:
        self.modulus = modulus
        self.exponent = exponent
        self._private = private
        self.size = (modulus.bit_length() + 7) // 8

    def sign(self, message: str) -> str:
        """EMSA-PKCS1-v1_5 over SHA-256, base64 — what `openssl_sign` produces."""
        digest = hashlib.sha256(message.encode("utf-8")).digest()
        tail = SHA256_DIGEST_INFO + digest
        encoded = b"\x00\x01" + b"\xff" * (self.size - len(tail) - 3) + b"\x00" + tail

        signature = pow(int.from_bytes(encoded, "big"), self._private, self.modulus)

        return base64.b64encode(signature.to_bytes(self.size, "big")).decode("ascii")

    @property
    def public_pem(self) -> str:
        """The public half, as a SubjectPublicKeyInfo PEM."""
        rsa = _tlv(0x30, _integer(self.modulus) + _integer(self.exponent))
        spki = _tlv(0x30, RSA_ALGORITHM_ID + _tlv(0x03, b"\x00" + rsa))
        body = base64.b64encode(spki).decode("ascii")
        wrapped = "\n".join(body[i : i + 64] for i in range(0, len(body), 64))

        return f"-----BEGIN PUBLIC KEY-----\n{wrapped}\n-----END PUBLIC KEY-----\n"


def generate(bits: int = 2048, seed: int | None = 20260729) -> TestKey:
    """A usable keypair.

    Seeded by default so a failure is reproducible — a test that fails once in
    fifty runs against a key nobody can regenerate is worse than no test.
    """
    rng = random.Random(seed)
    exponent = 65537

    while True:
        p = _prime(bits // 2, rng)
        q = _prime(bits // 2, rng)

        if p == q:
            continue

        modulus = p * q

        if modulus.bit_length() != bits:
            continue

        totient = (p - 1) * (q - 1)

        if math.gcd(exponent, totient) != 1:
            continue

        return TestKey(modulus, exponent, pow(exponent, -1, totient))


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


def _prime(bits: int, rng: random.Random) -> int:
    while True:
        candidate = rng.getrandbits(bits) | (1 << (bits - 1)) | (1 << (bits - 2)) | 1

        if any(candidate % small == 0 for small in SMALL_PRIMES):
            continue

        if _probably_prime(candidate, rng):
            return candidate


def _probably_prime(n: int, rng: random.Random, rounds: int = 20) -> bool:
    """Miller-Rabin."""
    d = n - 1
    r = 0

    while d % 2 == 0:
        d //= 2
        r += 1

    for _ in range(rounds):
        a = rng.randrange(2, n - 1)
        x = pow(a, d, n)

        if x in (1, n - 1):
            continue

        for _ in range(r - 1):
            x = pow(x, 2, n)

            if x == n - 1:
                break
        else:
            return False

    return True


def _tlv(tag: int, value: bytes) -> bytes:
    length = len(value)

    if length < 0x80:
        header = bytes([tag, length])
    else:
        encoded = length.to_bytes((length.bit_length() + 7) // 8, "big")
        header = bytes([tag, 0x80 | len(encoded)]) + encoded

    return header + value


def _integer(value: int) -> bytes:
    # One byte more than strictly needed whenever the top bit would be set,
    # because a DER INTEGER is signed and would otherwise read as negative.
    width = (value.bit_length() // 8) + 1

    return _tlv(0x02, value.to_bytes(width, "big"))
