package com.codebyte.api.support;

import java.io.StringWriter;
import java.math.BigInteger;
import java.security.KeyPair;
import java.security.KeyPairGenerator;
import java.security.PrivateKey;
import java.security.PublicKey;
import java.security.Security;
import java.security.cert.X509Certificate;
import java.time.Instant;
import java.time.temporal.ChronoUnit;
import java.util.List;
import java.util.concurrent.atomic.AtomicLong;
import org.bouncycastle.asn1.pkcs.PrivateKeyInfo;
import org.bouncycastle.asn1.x500.X500Name;
import org.bouncycastle.asn1.x509.BasicConstraints;
import org.bouncycastle.asn1.x509.Extension;
import org.bouncycastle.asn1.x509.GeneralName;
import org.bouncycastle.asn1.x509.GeneralNames;
import org.bouncycastle.cert.jcajce.JcaX509CertificateConverter;
import org.bouncycastle.cert.jcajce.JcaX509v3CertificateBuilder;
import org.bouncycastle.jce.provider.BouncyCastleProvider;
import org.bouncycastle.openssl.PKCS8Generator;
import org.bouncycastle.openssl.jcajce.JcaPEMWriter;
import org.bouncycastle.openssl.jcajce.JcaPKCS8Generator;
import org.bouncycastle.openssl.jcajce.JceOpenSSLPKCS8EncryptorBuilder;
import org.bouncycastle.operator.ContentSigner;
import org.bouncycastle.operator.OutputEncryptor;
import org.bouncycastle.operator.jcajce.JcaContentSignerBuilder;
import org.bouncycastle.util.io.pem.PemObject;

/**
 * Generates a 3-level test PKI (root → intermediate → leaf) and renders PEM in the formats the SSL
 * loader must accept. No files are committed; everything is produced in memory (or written to a
 * temp dir by the caller).
 */
public final class TestPki {

    static {
        if (Security.getProvider(BouncyCastleProvider.PROVIDER_NAME) == null) {
            Security.addProvider(new BouncyCastleProvider());
        }
    }

    private static final AtomicLong SERIAL = new AtomicLong(1000);

    public final X509Certificate root;
    public final X509Certificate intermediate;
    public final X509Certificate leaf;
    public final PrivateKey leafKey;

    private TestPki(
            X509Certificate root,
            X509Certificate intermediate,
            X509Certificate leaf,
            PrivateKey leafKey) {
        this.root = root;
        this.intermediate = intermediate;
        this.leaf = leaf;
        this.leafKey = leafKey;
    }

    public static TestPki threeLevel() {
        return build(
                Instant.now().minus(1, ChronoUnit.DAYS), Instant.now().plus(365, ChronoUnit.DAYS));
    }

    /** Same chain but the leaf is already expired. */
    public static TestPki expiredLeaf() {
        return build(
                Instant.now().minus(60, ChronoUnit.DAYS), Instant.now().minus(1, ChronoUnit.DAYS));
    }

    /** Same chain but the leaf expires soon (still valid) — for the near-expiry health check. */
    public static TestPki leafExpiringInDays(int days) {
        return build(
                Instant.now().minus(1, ChronoUnit.DAYS), Instant.now().plus(days, ChronoUnit.DAYS));
    }

    private static TestPki build(Instant leafNotBefore, Instant leafNotAfter) {
        try {
            Instant caFrom = Instant.now().minus(2, ChronoUnit.DAYS);
            Instant caTo = Instant.now().plus(3650, ChronoUnit.DAYS);

            KeyPair rootKp = rsa();
            X509Certificate rootCert =
                    cert(
                            "CN=Dev Root CA",
                            rootKp.getPublic(),
                            "CN=Dev Root CA",
                            rootKp.getPrivate(),
                            true,
                            caFrom,
                            caTo,
                            null);

            KeyPair intKp = rsa();
            X509Certificate intCert =
                    cert(
                            "CN=Dev Intermediate CA",
                            intKp.getPublic(),
                            "CN=Dev Root CA",
                            rootKp.getPrivate(),
                            true,
                            caFrom,
                            caTo,
                            null);

            KeyPair leafKp = rsa();
            X509Certificate leafCert =
                    cert(
                            "CN=localhost",
                            leafKp.getPublic(),
                            "CN=Dev Intermediate CA",
                            intKp.getPrivate(),
                            false,
                            leafNotBefore,
                            leafNotAfter,
                            List.of(
                                    new GeneralName(GeneralName.dNSName, "localhost"),
                                    new GeneralName(GeneralName.iPAddress, "127.0.0.1")));

            return new TestPki(rootCert, intCert, leafCert, leafKp.getPrivate());
        } catch (Exception e) {
            throw new IllegalStateException("failed to build test PKI", e);
        }
    }

    private static KeyPair rsa() throws Exception {
        KeyPairGenerator kpg = KeyPairGenerator.getInstance("RSA");
        kpg.initialize(2048);
        return kpg.generateKeyPair();
    }

    private static X509Certificate cert(
            String subject,
            PublicKey subjectKey,
            String issuer,
            PrivateKey issuerKey,
            boolean ca,
            Instant notBefore,
            Instant notAfter,
            List<GeneralName> sans)
            throws Exception {
        JcaX509v3CertificateBuilder builder =
                new JcaX509v3CertificateBuilder(
                        new X500Name(issuer),
                        BigInteger.valueOf(SERIAL.incrementAndGet()),
                        java.util.Date.from(notBefore),
                        java.util.Date.from(notAfter),
                        new X500Name(subject),
                        subjectKey);
        builder.addExtension(Extension.basicConstraints, true, new BasicConstraints(ca));
        if (sans != null) {
            builder.addExtension(
                    Extension.subjectAlternativeName,
                    false,
                    new GeneralNames(sans.toArray(new GeneralName[0])));
        }
        ContentSigner signer = new JcaContentSignerBuilder("SHA256withRSA").build(issuerKey);
        return new JcaX509CertificateConverter().getCertificate(builder.build(signer));
    }

    // --- PEM rendering ---

    public String leafCertPem() {
        return pem(leaf);
    }

    /** CA file: intermediate then root (root will be stripped by the loader). */
    public String caChainPem() {
        return pem(intermediate) + pem(root);
    }

    /** CA file with the certificates in the wrong order (root first) — the loader must reorder. */
    public String caChainPemWrongOrder() {
        return pem(root) + pem(intermediate);
    }

    public String keyPkcs8Pem() {
        return pem(leafKey);
    }

    public String keyPkcs1Pem() {
        try {
            PrivateKeyInfo info = PrivateKeyInfo.getInstance(leafKey.getEncoded());
            byte[] pkcs1 = info.parsePrivateKey().toASN1Primitive().getEncoded();
            return pem(new PemObject("RSA PRIVATE KEY", pkcs1));
        } catch (Exception e) {
            throw new IllegalStateException(e);
        }
    }

    public String keyPkcs8EncryptedPem(String password) {
        try {
            OutputEncryptor encryptor =
                    new JceOpenSSLPKCS8EncryptorBuilder(PKCS8Generator.AES_256_CBC)
                            .setProvider(BouncyCastleProvider.PROVIDER_NAME)
                            .setPassword(password.toCharArray())
                            .build();
            PemObject pem = new JcaPKCS8Generator(leafKey, encryptor).generate();
            return pem(pem);
        } catch (Exception e) {
            throw new IllegalStateException(e);
        }
    }

    private static String pem(Object object) {
        StringWriter writer = new StringWriter();
        try (JcaPEMWriter pemWriter = new JcaPEMWriter(writer)) {
            pemWriter.writeObject(object);
        } catch (Exception e) {
            throw new IllegalStateException(e);
        }
        return writer.toString();
    }
}
