package com.codebyte.api.config.ssl;

import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.io.StringReader;
import java.nio.charset.StandardCharsets;
import java.nio.file.Files;
import java.nio.file.Path;
import java.security.PrivateKey;
import java.security.Security;
import java.security.Signature;
import java.security.cert.CertificateExpiredException;
import java.security.cert.CertificateFactory;
import java.security.cert.CertificateNotYetValidException;
import java.security.cert.X509Certificate;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
import javax.security.auth.x500.X500Principal;
import org.bouncycastle.asn1.pkcs.PrivateKeyInfo;
import org.bouncycastle.jce.provider.BouncyCastleProvider;
import org.bouncycastle.openssl.PEMDecryptorProvider;
import org.bouncycastle.openssl.PEMEncryptedKeyPair;
import org.bouncycastle.openssl.PEMKeyPair;
import org.bouncycastle.openssl.PEMParser;
import org.bouncycastle.openssl.jcajce.JcaPEMKeyConverter;
import org.bouncycastle.openssl.jcajce.JceOpenSSLPKCS8DecryptorProviderBuilder;
import org.bouncycastle.openssl.jcajce.JcePEMDecryptorProviderBuilder;
import org.bouncycastle.operator.InputDecryptorProvider;
import org.bouncycastle.pkcs.PKCS8EncryptedPrivateKeyInfo;

/**
 * Reads the cert/CA/key trio, assembles the certificate chain in memory (leaf → intermediates, root
 * stripped), loads the private key (PKCS#1 or PKCS#8, plain or encrypted), and validates the lot.
 * Every failure mode throws a {@link SslConfigurationException} with a specific, actionable
 * message.
 */
public class PemSslLoader {

    private static final byte[] MATCH_PROBE =
            "key-cert-match-probe".getBytes(StandardCharsets.UTF_8);
    private static final int MAX_CHAIN = 16;

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

    private final SslProperties properties;

    public PemSslLoader(SslProperties properties) {
        this.properties = properties;
    }

    /** Loads and fully validates the trio. */
    public PemSslMaterial load() {
        Path certPath = requireConfigured(properties.certificate(), "SSL_CERT_PATH");
        Path caPath = requireConfigured(properties.ca(), "SSL_CA_PATH");
        Path keyPath = requireConfigured(properties.privateKey(), "SSL_KEY_PATH");

        List<X509Certificate> certFile =
                parseCertificates(readReadable(certPath, "certificate (SSL_CERT_PATH)"));
        List<X509Certificate> caFile =
                parseCertificates(readReadable(caPath, "CA chain (SSL_CA_PATH)"));
        List<X509Certificate> chain = assembleChain(certFile, caFile);

        char[] password = passwordChars();
        PrivateKey privateKey =
                parsePrivateKey(readReadable(keyPath, "private key (SSL_KEY_PATH)"), password);

        validateKeyMatchesLeaf(privateKey, chain.get(0));
        validateChain(chain);
        return new PemSslMaterial(privateKey, chain);
    }

    // --- parsing ---

    static List<X509Certificate> parseCertificates(String pem) {
        try {
            CertificateFactory factory = CertificateFactory.getInstance("X.509");
            var certs =
                    factory.generateCertificates(
                            new ByteArrayInputStream(pem.getBytes(StandardCharsets.UTF_8)));
            List<X509Certificate> result = new ArrayList<>();
            certs.forEach(c -> result.add((X509Certificate) c));
            if (result.isEmpty()) {
                throw new SslConfigurationException("no X.509 certificates found in PEM content");
            }
            return result;
        } catch (SslConfigurationException e) {
            throw e;
        } catch (Exception e) {
            throw new SslConfigurationException(
                    "failed to parse X.509 certificates: " + e.getMessage(), e);
        }
    }

    static PrivateKey parsePrivateKey(String pem, char[] password) {
        try (PEMParser parser = new PEMParser(new StringReader(pem))) {
            Object object = parser.readObject();
            JcaPEMKeyConverter converter =
                    new JcaPEMKeyConverter().setProvider(BouncyCastleProvider.PROVIDER_NAME);
            if (object == null) {
                throw new SslConfigurationException("no PEM object found in the private key file");
            }
            return switch (object) {
                case PEMEncryptedKeyPair encrypted -> { // PKCS#1 encrypted
                    requirePassword(password);
                    PEMDecryptorProvider decryptor =
                            new JcePEMDecryptorProviderBuilder().build(password);
                    yield converter.getKeyPair(encrypted.decryptKeyPair(decryptor)).getPrivate();
                }
                case PKCS8EncryptedPrivateKeyInfo encrypted -> { // PKCS#8 encrypted
                    requirePassword(password);
                    InputDecryptorProvider decryptor =
                            new JceOpenSSLPKCS8DecryptorProviderBuilder()
                                    .setProvider(BouncyCastleProvider.PROVIDER_NAME)
                                    .build(password);
                    yield converter.getPrivateKey(encrypted.decryptPrivateKeyInfo(decryptor));
                }
                case PEMKeyPair keyPair -> converter.getKeyPair(keyPair).getPrivate(); // PKCS#1
                case PrivateKeyInfo info -> converter.getPrivateKey(info); // PKCS#8
                default ->
                        throw new SslConfigurationException(
                                "unsupported private key format: "
                                        + object.getClass().getSimpleName());
            };
        } catch (SslConfigurationException e) {
            throw e;
        } catch (Exception e) {
            throw new SslConfigurationException(
                    "failed to load the private key (wrong passphrase, or unsupported/corrupt key):"
                            + " "
                            + e.getMessage(),
                    e);
        }
    }

    // --- chain assembly ---

    static List<X509Certificate> assembleChain(
            List<X509Certificate> certFile, List<X509Certificate> caFile) {
        X509Certificate leaf = certFile.get(0);
        Map<X500Principal, X509Certificate> bySubject = new HashMap<>();
        for (X509Certificate cert : certFile) {
            bySubject.putIfAbsent(cert.getSubjectX500Principal(), cert);
        }
        for (X509Certificate cert : caFile) {
            bySubject.putIfAbsent(cert.getSubjectX500Principal(), cert);
        }

        List<X509Certificate> chain = new ArrayList<>();
        chain.add(leaf);
        Set<X500Principal> seen = new HashSet<>();
        seen.add(leaf.getSubjectX500Principal());
        X509Certificate current = leaf;
        while (!isSelfSigned(current) && chain.size() < MAX_CHAIN) {
            X509Certificate issuer = bySubject.get(current.getIssuerX500Principal());
            if (issuer == null || isSelfSigned(issuer)) {
                break; // top of provided material, or self-signed root (stripped)
            }
            if (!seen.add(issuer.getSubjectX500Principal())) {
                break; // loop guard
            }
            chain.add(issuer);
            current = issuer;
        }
        return chain;
    }

    // --- validation ---

    static void validateKeyMatchesLeaf(PrivateKey privateKey, X509Certificate leaf) {
        String algorithm =
                switch (privateKey.getAlgorithm()) {
                    case "RSA" -> "SHA256withRSA";
                    case "EC", "ECDSA" -> "SHA256withECDSA";
                    default ->
                            throw new SslConfigurationException(
                                    "unsupported key algorithm: " + privateKey.getAlgorithm());
                };
        try {
            Signature signer = Signature.getInstance(algorithm);
            signer.initSign(privateKey);
            signer.update(MATCH_PROBE);
            byte[] signature = signer.sign();

            Signature verifier = Signature.getInstance(algorithm);
            verifier.initVerify(leaf.getPublicKey());
            verifier.update(MATCH_PROBE);
            if (!verifier.verify(signature)) {
                throw new SslConfigurationException(
                        "the private key does not match the leaf certificate's public key");
            }
        } catch (SslConfigurationException e) {
            throw e;
        } catch (Exception e) {
            throw new SslConfigurationException(
                    "the private key does not match the leaf certificate: " + e.getMessage(), e);
        }
    }

    static void validateChain(List<X509Certificate> chain) {
        for (int i = 0; i < chain.size() - 1; i++) {
            X509Certificate cert = chain.get(i);
            X509Certificate issuer = chain.get(i + 1);
            if (!cert.getIssuerX500Principal().equals(issuer.getSubjectX500Principal())) {
                throw new SslConfigurationException(
                        "certificate chain is broken: issuer of '"
                                + cert.getSubjectX500Principal()
                                + "' does not match subject of the next certificate");
            }
            try {
                cert.verify(issuer.getPublicKey());
            } catch (Exception e) {
                throw new SslConfigurationException(
                        "certificate chain signature is invalid at position "
                                + i
                                + ": "
                                + e.getMessage(),
                        e);
            }
        }
        X509Certificate leaf = chain.get(0);
        try {
            leaf.checkValidity();
        } catch (CertificateExpiredException e) {
            throw new SslConfigurationException(
                    "the leaf certificate expired on " + leaf.getNotAfter(), e);
        } catch (CertificateNotYetValidException e) {
            throw new SslConfigurationException(
                    "the leaf certificate is not valid until " + leaf.getNotBefore(), e);
        }
    }

    private static boolean isSelfSigned(X509Certificate cert) {
        return cert.getSubjectX500Principal().equals(cert.getIssuerX500Principal());
    }

    // --- file handling ---

    private static Path requireConfigured(String value, String envVar) {
        if (value == null || value.isBlank()) {
            throw new SslConfigurationException(envVar + " is not set but SSL is enabled");
        }
        String path = value.startsWith("file:") ? value.substring("file:".length()) : value;
        return Path.of(path);
    }

    private static String readReadable(Path path, String label) {
        if (!Files.exists(path)) {
            throw new SslConfigurationException(label + " file does not exist: " + path);
        }
        if (!Files.isRegularFile(path)) {
            throw new SslConfigurationException(label + " is not a regular file: " + path);
        }
        if (!Files.isReadable(path)) {
            throw new SslConfigurationException(
                    label
                            + " is not readable by the application user ("
                            + currentUser()
                            + "): "
                            + path
                            + " — this is the most common deployment failure; fix"
                            + " ownership/permissions so this uid/gid can read the file (or run the"
                            + " container with a matching user).");
        }
        try {
            return Files.readString(path, StandardCharsets.UTF_8);
        } catch (IOException e) {
            throw new SslConfigurationException(label + " could not be read: " + path, e);
        }
    }

    private char[] passwordChars() {
        String password = properties.privateKeyPassword();
        return (password == null || password.isEmpty()) ? null : password.toCharArray();
    }

    private static void requirePassword(char[] password) {
        if (password == null) {
            throw new SslConfigurationException(
                    "the private key is encrypted but SSL_KEY_PASSWORD is not set");
        }
    }

    private static String currentUser() {
        try {
            var unix = new com.sun.security.auth.module.UnixSystem();
            return "uid=" + unix.getUid() + " gid=" + unix.getGid();
        } catch (Throwable ignored) {
            return "user=" + System.getProperty("user.name", "unknown");
        }
    }
}
