from __future__ import annotations

"""Tools for generating RSA key pairs for licence signing."""

from dataclasses import dataclass
from pathlib import Path

from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa


@dataclass
class LicenceKeyPair:
    private_key_pem: bytes
    public_key_pem: bytes


def generate_key_pair() -> LicenceKeyPair:
    private_key = rsa.generate_private_key(public_exponent=65537, key_size=4096)
    private_pem = private_key.private_bytes(
        encoding=serialization.Encoding.PEM,
        format=serialization.PrivateFormat.PKCS8,
        encryption_algorithm=serialization.NoEncryption(),
    )
    public_pem = private_key.public_key().public_bytes(
        encoding=serialization.Encoding.PEM,
        format=serialization.PublicFormat.SubjectPublicKeyInfo,
    )
    return LicenceKeyPair(private_pem, public_pem)


def write_key_pair(output_dir: Path) -> LicenceKeyPair:
    output_dir.mkdir(parents=True, exist_ok=True)
    key_pair = generate_key_pair()
    (output_dir / "license_private.pem").write_bytes(key_pair.private_key_pem)
    (output_dir / "license_public.pem").write_bytes(key_pair.public_key_pem)
    return key_pair
