#!/usr/bin/python3
"""Drop-in replacement for Dogtag's PKCS12Export Java tool.

Exports all certificates and private keys from an NSS database into
a single PKCS#12 file.  Used by FreeIPA's custodiainstance._get_keys()
during replica key transfer.

Usage:
    PKCS12Export -d <nssdb_dir> -p <nssdb_pwdfile> \
                 -w <export_pwdfile> -o <output.p12>

Strategy:
  1. List nicknames via certutil
  2. For each nickname, export cert DER via certutil
  3. For each nickname with a private key, export key via pk12util
     then openssl, then encrypt as PKCS8ShroudedKeyBag via openssl pkcs8
  4. Build SafeBags (cert bags + shrouded key bags) and assemble the PFX
     with HMAC-SHA1 MAC

This works for all key types NSS supports (RSA, EC, ML-DSA) because
cert export uses certutil directly and key export goes through openssl
which handles PQC key types via provider params.
"""

import argparse
import hashlib
import hmac as hmac_mod
import math
import os
import shutil
import subprocess
import sys
import tempfile


# -- minimal DER helpers --------------------------------------------------

def _der_len(n):
    if n < 0x80:
        return bytes([n])
    if n < 0x100:
        return bytes([0x81, n])
    if n < 0x10000:
        return bytes([0x82, (n >> 8) & 0xff, n & 0xff])
    return bytes([0x83, (n >> 16) & 0xff, (n >> 8) & 0xff, n & 0xff])


def _der(tag, value):
    return bytes([tag]) + _der_len(len(value)) + value


def _seq(*parts):
    return _der(0x30, b"".join(parts))


def _set(*parts):
    return _der(0x31, b"".join(parts))


def _octet(data):
    return _der(0x04, data)


def _integer(v):
    if v < 0x80:
        return _der(0x02, bytes([v]))
    bs = v.to_bytes((v.bit_length() + 8) // 8, "big")
    return _der(0x02, bs)


def _oid(dotted):
    parts = [int(x) for x in dotted.split(".")]
    enc = [40 * parts[0] + parts[1]]
    for p in parts[2:]:
        if p < 0x80:
            enc.append(p)
        else:
            chunks = []
            while p:
                chunks.append(p & 0x7f)
                p >>= 7
            chunks.reverse()
            for j, c in enumerate(chunks):
                enc.append(c | 0x80 if j < len(chunks) - 1 else c)
    return _der(0x06, bytes(enc))


def _explicit(tag_num, content):
    return bytes([0xa0 | tag_num]) + _der_len(len(content)) + content


def _bmpstring(text):
    return _der(0x1e, text.encode("utf-16-be"))


def _content_info(oid_str, content_bytes):
    return _seq(_oid(oid_str), _explicit(0, content_bytes))


# -- PKCS#12 OIDs ---------------------------------------------------------

_OID_DATA = "1.2.840.113549.1.7.1"
_OID_CERT_BAG = "1.2.840.113549.1.12.10.1.3"
_OID_SHROUDED_KEY_BAG = "1.2.840.113549.1.12.10.1.2"
_OID_X509 = "1.2.840.113549.1.9.22.1"
_OID_FNAME = "1.2.840.113549.1.9.20"
_OID_LKID = "1.2.840.113549.1.9.21"
_OID_SHA1 = "1.3.14.3.2.26"


# -- PKCS#12 KDF (RFC 7292 Appendix B) ------------------------------------

def _pkcs12_kdf(password_bmp, salt, iterations, n, id_byte):
    """Derive key material from password using PKCS#12 KDF."""
    v = 64   # SHA-1 block size
    u = 20   # SHA-1 digest size

    D = bytes([id_byte]) * v

    if salt:
        s_len = v * math.ceil(len(salt) / v)
        S = (salt * (s_len // len(salt) + 1))[:s_len]
    else:
        S = b""

    if password_bmp:
        p_len = v * math.ceil(len(password_bmp) / v)
        P = (password_bmp * (p_len // len(password_bmp) + 1))[:p_len]
    else:
        P = b""

    I = bytearray(S + P)  # noqa: E741 — RFC 7292 Appendix B name
    c = math.ceil(n / u)
    result = b""

    for i in range(c):
        A = hashlib.sha1(D + bytes(I)).digest()
        for _ in range(1, iterations):
            A = hashlib.sha1(A).digest()
        result += A

        if i < c - 1:
            B = (A * (v // u + 1))[:v]
            for j in range(len(I) // v):
                carry = 1
                for k in range(v - 1, -1, -1):
                    t = I[j * v + k] + B[k] + carry
                    I[j * v + k] = t & 0xff
                    carry = t >> 8

    return result[:n]


def _compute_mac(password, data, iterations=2048):
    """Compute PKCS#12 MAC (HMAC-SHA1) over authSafe content."""
    salt = os.urandom(8)
    password_bmp = password.encode("utf-16-be") + b"\x00\x00"
    mac_key = _pkcs12_kdf(password_bmp, salt, iterations, 20, 3)
    mac_value = hmac_mod.new(mac_key, data, hashlib.sha1).digest()

    digest_info = _seq(
        _seq(_oid(_OID_SHA1), b"\x05\x00"),
        _octet(mac_value),
    )
    return _seq(digest_info, _octet(salt), _integer(iterations))


# -- bag builders ----------------------------------------------------------

def _attrs(name, key_id):
    return _set(
        _seq(_oid(_OID_FNAME), _set(_bmpstring(name))),
        _seq(_oid(_OID_LKID), _set(_octet(key_id))),
    )


def _cert_bag(cert_der, name, key_id):
    value = _seq(_oid(_OID_X509), _explicit(0, _octet(cert_der)))
    return _seq(_oid(_OID_CERT_BAG), _explicit(0, value), _attrs(name, key_id))


def _shrouded_key_bag(encrypted_pkcs8_der, name, key_id):
    return _seq(
        _oid(_OID_SHROUDED_KEY_BAG),
        _explicit(0, encrypted_pkcs8_der),
        _attrs(name, key_id),
    )


# -- export helpers --------------------------------------------------------

def _export_cert_der(dbdir, nickname):
    """Export certificate as raw DER via certutil."""
    result = subprocess.run(
        ["certutil", "-L", "-d", dbdir, "-n", nickname, "-r"],
        capture_output=True, check=False,
    )
    if result.returncode != 0 or not result.stdout:
        return None
    return result.stdout


def _export_encrypted_key_der(dbdir, nickname, pwdfile, export_pwdfile):
    """Export private key as EncryptedPrivateKeyInfo DER.

    Returns None if the nickname has no private key.
    Uses pk12util to get the key out of NSSDB, openssl pkcs12 to
    extract the PEM, and openssl pkcs8 to re-encrypt as PKCS#12 PBE.
    """
    tmpdir = tempfile.mkdtemp()
    try:
        p12_path = os.path.join(tmpdir, "temp.p12")

        # Step 1: pk12util exports cert+key to a temp PKCS#12
        result = subprocess.run(
            ["pk12util", "-o", p12_path, "-d", dbdir,
             "-n", nickname, "-k", pwdfile, "-w", export_pwdfile],
            capture_output=True, check=False,
        )
        if result.returncode != 0:
            return None  # no private key for this nickname

        # Step 2: openssl pkcs12 extracts the key PEM
        key_pem = _openssl_extract_key_pem(p12_path, export_pwdfile)
        if key_pem is None:
            return None

        # Step 3: openssl pkcs8 re-encrypts as EncryptedPrivateKeyInfo
        return _openssl_encrypt_pkcs8(key_pem, export_pwdfile)

    finally:
        shutil.rmtree(tmpdir, ignore_errors=True)


def _openssl_extract_key_pem(p12_path, pwdfile):
    """Extract private key PEM from a PKCS#12 file via openssl."""
    # Try without PQC provider param first, then with it
    for extra_args in [
        [],
        ["-legacy"],
        ["-provparam", "ml-dsa.output_formats=seed-only"],
        ["-legacy", "-provparam", "ml-dsa.output_formats=seed-only"],
    ]:
        cmd = [
            "openssl", "pkcs12",
            "-in", p12_path,
            "-nocerts", "-nodes",
            "-passin", f"file:{pwdfile}",
        ] + extra_args
        result = subprocess.run(cmd, capture_output=True, check=False)
        if result.returncode == 0 and b"PRIVATE KEY" in result.stdout:
            return result.stdout
    return None


def _openssl_encrypt_pkcs8(key_pem, pwdfile):
    """Encrypt a PEM private key as EncryptedPrivateKeyInfo DER.

    Uses PBE-SHA1-3DES (PKCS#12 PBE) for compatibility with pk12util.
    """
    for extra_args in [
        [],
        ["-provparam", "ml-dsa.output_formats=seed-only"],
    ]:
        cmd = [
            "openssl", "pkcs8", "-topk8",
            "-v1", "PBE-SHA1-3DES",
            "-in", "/dev/stdin",
            "-outform", "DER",
            "-passout", f"file:{pwdfile}",
        ] + extra_args
        result = subprocess.run(
            cmd, input=key_pem, capture_output=True, check=False,
        )
        if result.returncode == 0 and result.stdout:
            return result.stdout
    return None


# -- main ------------------------------------------------------------------

def list_nicknames(dbdir):
    result = subprocess.run(
        ["certutil", "-L", "-d", dbdir],
        capture_output=True, text=True, check=False,
    )
    if result.returncode != 0:
        return []
    nicks = []
    in_header = True
    for line in result.stdout.splitlines():
        if not line.strip():
            in_header = False
            continue
        if in_header:
            continue
        parts = line.rsplit(None, 1)
        if len(parts) == 2 and parts[0].strip():
            nicks.append(parts[0].strip())
    return nicks


def main():
    parser = argparse.ArgumentParser(description="Export NSSDB to PKCS#12")
    parser.add_argument("-d", dest="dbdir", required=True)
    parser.add_argument("-p", dest="pwdfile", required=True)
    parser.add_argument("-w", dest="export_pwdfile", required=True)
    parser.add_argument("-o", dest="output", required=True)
    args = parser.parse_args()

    with open(args.export_pwdfile) as f:
        export_pwd = f.read().strip()

    nicknames = list_nicknames(args.dbdir)
    if not nicknames:
        print("No certificates found in NSSDB", file=sys.stderr)
        sys.exit(1)

    bags = []
    for nickname in nicknames:
        # Export certificate DER
        cert_der = _export_cert_der(args.dbdir, nickname)
        if cert_der is None:
            continue

        kid = hashlib.sha1(cert_der).digest()
        bags.append(_cert_bag(cert_der, nickname, kid))

        # Export private key (if present)
        enc_key_der = _export_encrypted_key_der(
            args.dbdir, nickname, args.pwdfile, args.export_pwdfile,
        )
        if enc_key_der is not None:
            bags.append(_shrouded_key_bag(enc_key_der, nickname, kid))

    if not bags:
        print("Failed to export any certificates", file=sys.stderr)
        sys.exit(1)

    safe_contents = _seq(*bags)
    auth_safe = _seq(_content_info(_OID_DATA, _octet(safe_contents)))
    mac_data = _compute_mac(export_pwd, auth_safe)

    pfx = _seq(
        _integer(3),
        _content_info(_OID_DATA, _octet(auth_safe)),
        mac_data,
    )

    with open(args.output, "wb") as f:
        f.write(pfx)


if __name__ == "__main__":
    main()
