#!/usr/bin/env python3
"""
Cert-creation primitives for the iOS native TestFlight action.

Owns the cryptographic legwork (RSA key + CSR), Apple's per-team cert
cap refusal (never automatic revocation), and PKCS12 serialisation. Split
from ``prepare_signing.py`` so the orchestrator there stays focused on
the load-or-regen control flow and the file count stays under the
project's per-file limits.

Nothing here touches the on-disk cache — callers receive raw bytes and
hand them to ``creds_store`` for atomic persistence.
"""

from __future__ import annotations

import base64
import os

from asc_common import request
import certificate_failure
from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.hazmat.primitives.serialization import pkcs12
from cryptography.x509.oid import NameOID


def certificate_cap_policy() -> str:
    """Refuse destructive legacy policy before any provider call."""
    policy = os.getenv("CERTIFICATE_CAP_POLICY", "fail").strip().lower()
    if policy != "fail":
        raise SystemExit("invalid CERTIFICATE_CAP_POLICY; use fail. Automatic certificate revocation is disabled")
    return policy


def generate_key_and_csr() -> tuple[rsa.RSAPrivateKey, bytes]:
    """Generate a fresh 2048-bit RSA key and complete PEM CSR.

    Returns the private key (kept in memory; the caller serialises it
    into the PKCS12) and the PEM CSR, including its framing and newlines,
    for Apple's ``csrContent`` field.
    """
    private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
    csr = (
        x509.CertificateSigningRequestBuilder()
        .subject_name(
            x509.Name(
                [
                    x509.NameAttribute(NameOID.COMMON_NAME, "Daemux CI"),
                    x509.NameAttribute(NameOID.COUNTRY_NAME, "US"),
                ]
            )
        )
        .sign(private_key, hashes.SHA256())
    )
    csr_pem = csr.public_bytes(serialization.Encoding.PEM)
    return private_key, csr_pem


def create_distribution_cert(token: str, csr_pem: str) -> tuple[str, bytes]:
    """POST one CSR; a full team cap preserves every existing certificate."""
    certificate_cap_policy()
    body = {
        "data": {
            "type": "certificates",
            "attributes": {
                "csrContent": csr_pem,
                "certificateType": "DISTRIBUTION",
            },
        }
    }
    resp = request(
        "POST", "/certificates", token, json_body=body, allow_status={409}
    )
    if resp.status_code == 409:
        certificate_failure.refuse(resp)
    data = resp.json()["data"]
    cert_id = data["id"]
    cert_der = base64.b64decode(data["attributes"]["certificateContent"])
    print(f"Created DISTRIBUTION cert {cert_id}")
    return cert_id, cert_der


def serialize_p12(
    private_key: rsa.RSAPrivateKey, cert_der: bytes, passwd: str
) -> bytes:
    """Return the PKCS12 bytes encrypted with ``passwd``.

    Caller decides whether to write to disk via the cache layer or as a
    transient artifact in ``$RUNNER_TEMP``.
    """
    cert = x509.load_der_x509_certificate(cert_der)
    return pkcs12.serialize_key_and_certificates(
        name=b"Daemux CI",
        key=private_key,
        cert=cert,
        cas=None,
        encryption_algorithm=serialization.BestAvailableEncryption(
            passwd.encode()
        ),
    )
