diff --git a/src/server_enrollment.py b/src/server_enrollment.py new file mode 100644 index 0000000..2ac60a0 --- /dev/null +++ b/src/server_enrollment.py @@ -0,0 +1,267 @@ +import os +import datetime +import ipaddress +from typing import Tuple, List, Optional +from cryptography import x509 +from cryptography.x509.oid import NameOID, ExtendedKeyUsageOID +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import rsa + + +def calculate_cert_fingerprint(cert_pem: str) -> str: + """Computes SHA-256 fingerprint for a PEM-encoded X.509 certificate.""" + cert = x509.load_pem_x509_certificate(cert_pem.encode("utf-8")) + return cert.fingerprint(hashes.SHA256()).hex().upper() + + +def generate_ca_if_needed(cert_dir: str = "certs", common_name: str = "LOGAR-Root-CA") -> Tuple[x509.Certificate, rsa.RSAPrivateKey, str, str]: + """ + Loads an existing Root CA or generates a self-signed Root CA certificate and private key. + Returns (ca_cert_obj, ca_key_obj, ca_cert_pem, ca_key_pem). + """ + os.makedirs(cert_dir, exist_ok=True) + ca_cert_path = os.path.join(cert_dir, "ca.crt") + ca_key_path = os.path.join(cert_dir, "ca.key") + + if os.path.exists(ca_cert_path) and os.path.exists(ca_key_path): + with open(ca_cert_path, "r", encoding="utf-8") as f: + ca_cert_pem = f.read() + with open(ca_key_path, "r", encoding="utf-8") as f: + ca_key_pem = f.read() + ca_cert = x509.load_pem_x509_certificate(ca_cert_pem.encode("utf-8")) + ca_key = serialization.load_pem_private_key(ca_key_pem.encode("utf-8"), password=None) + return ca_cert, ca_key, ca_cert_pem, ca_key_pem + + # Generate RSA 4096 private key for Root CA + ca_key = rsa.generate_private_key(public_exponent=65537, key_size=4096) + subject = issuer = x509.Name([ + x509.NameAttribute(NameOID.COUNTRY_NAME, "AT"), + x509.NameAttribute(NameOID.ORGANIZATION_NAME, "LOGAR"), + x509.NameAttribute(NameOID.COMMON_NAME, common_name), + ]) + + now = datetime.datetime.now(datetime.timezone.utc) + ca_cert = ( + x509.CertificateBuilder() + .subject_name(subject) + .issuer_name(issuer) + .public_key(ca_key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - datetime.timedelta(minutes=5)) + .not_valid_after(now + datetime.timedelta(days=3650)) + .add_extension(x509.BasicConstraints(ca=True, path_length=None), critical=True) + .add_extension( + x509.KeyUsage( + digital_signature=True, + key_encipherment=False, + key_cert_sign=True, + crl_sign=True, + content_commitment=False, + data_encipherment=False, + key_agreement=False, + encipher_only=False, + decipher_only=False + ), + critical=True + ) + .add_extension( + x509.SubjectKeyIdentifier.from_public_key(ca_key.public_key()), + critical=False + ) + .sign(ca_key, hashes.SHA256()) + ) + + ca_cert_pem = ca_cert.public_bytes(serialization.Encoding.PEM).decode("utf-8") + ca_key_pem = ca_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.TraditionalOpenSSL, + encryption_algorithm=serialization.NoEncryption() + ).decode("utf-8") + + with open(ca_cert_path, "w", encoding="utf-8") as f: + f.write(ca_cert_pem) + with open(ca_key_path, "w", encoding="utf-8") as f: + f.write(ca_key_pem) + + try: + os.chmod(ca_key_path, 0o600) + except Exception: + pass + + return ca_cert, ca_key, ca_cert_pem, ca_key_pem + + +def generate_server_cert_if_needed( + ca_cert: x509.Certificate, + ca_key: rsa.RSAPrivateKey, + hostnames: Optional[List[str]] = None, + cert_dir: str = "certs", + days_valid: int = 825 +) -> Tuple[x509.Certificate, rsa.RSAPrivateKey, str, str]: + """ + Loads an existing server certificate or generates a new server TLS certificate signed by the Root CA. + Includes SANs for localhost, 127.0.0.1, and specified hostnames. + """ + os.makedirs(cert_dir, exist_ok=True) + server_cert_path = os.path.join(cert_dir, "server.crt") + server_key_path = os.path.join(cert_dir, "server.key") + + if os.path.exists(server_cert_path) and os.path.exists(server_key_path): + with open(server_cert_path, "r", encoding="utf-8") as f: + server_cert_pem = f.read() + with open(server_key_path, "r", encoding="utf-8") as f: + server_key_pem = f.read() + srv_cert = x509.load_pem_x509_certificate(server_cert_pem.encode("utf-8")) + srv_key = serialization.load_pem_private_key(server_key_pem.encode("utf-8"), password=None) + return srv_cert, srv_key, server_cert_pem, server_key_pem + + server_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + subject = x509.Name([ + x509.NameAttribute(NameOID.COUNTRY_NAME, "AT"), + x509.NameAttribute(NameOID.ORGANIZATION_NAME, "LOGAR"), + x509.NameAttribute(NameOID.COMMON_NAME, "LOGAR-Server-Hub"), + ]) + + san_list = [ + x509.DNSName("localhost"), + x509.DNSName("LOGAR-Server-Hub"), + x509.IPAddress(ipaddress.IPv4Address("127.0.0.1")), + x509.IPAddress(ipaddress.IPv6Address("::1")), + ] + + if hostnames: + for host in hostnames: + if not host: + continue + try: + ip_obj = ipaddress.ip_address(host) + san_list.append(x509.IPAddress(ip_obj)) + except ValueError: + san_list.append(x509.DNSName(host)) + + now = datetime.datetime.now(datetime.timezone.utc) + server_cert = ( + x509.CertificateBuilder() + .subject_name(subject) + .issuer_name(ca_cert.subject) + .public_key(server_key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - datetime.timedelta(minutes=5)) + .not_valid_after(now + datetime.timedelta(days=days_valid)) + .add_extension(x509.BasicConstraints(ca=False, path_length=None), critical=True) + .add_extension( + x509.KeyUsage( + digital_signature=True, + key_encipherment=True, + key_cert_sign=False, + crl_sign=False, + content_commitment=False, + data_encipherment=False, + key_agreement=False, + encipher_only=False, + decipher_only=False + ), + critical=True + ) + .add_extension( + x509.ExtendedKeyUsage([ExtendedKeyUsageOID.SERVER_AUTH]), + critical=False + ) + .add_extension( + x509.SubjectKeyIdentifier.from_public_key(server_key.public_key()), + critical=False + ) + .add_extension( + x509.AuthorityKeyIdentifier.from_issuer_public_key(ca_key.public_key()), + critical=False + ) + .add_extension(x509.SubjectAlternativeName(san_list), critical=False) + .sign(ca_key, hashes.SHA256()) + ) + + server_cert_pem = server_cert.public_bytes(serialization.Encoding.PEM).decode("utf-8") + server_key_pem = server_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.TraditionalOpenSSL, + encryption_algorithm=serialization.NoEncryption() + ).decode("utf-8") + + with open(server_cert_path, "w", encoding="utf-8") as f: + f.write(server_cert_pem) + with open(server_key_path, "w", encoding="utf-8") as f: + f.write(server_key_pem) + + try: + os.chmod(server_key_path, 0o600) + except Exception: + pass + + return server_cert, server_key, server_cert_pem, server_key_pem + + +def issue_client_cert( + client_id: str, + ca_cert: x509.Certificate, + ca_key: rsa.RSAPrivateKey, + days_valid: int = 365 +) -> Tuple[str, str]: + """ + Generates a 2048-bit RSA private key and signs an X.509 client certificate + with Common Name set to client_id. + Returns (cert_pem, key_pem). + """ + client_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + + subject = x509.Name([ + x509.NameAttribute(NameOID.COUNTRY_NAME, "AT"), + x509.NameAttribute(NameOID.ORGANIZATION_NAME, "LOGAR"), + x509.NameAttribute(NameOID.COMMON_NAME, client_id), + ]) + + now = datetime.datetime.now(datetime.timezone.utc) + cert = ( + x509.CertificateBuilder() + .subject_name(subject) + .issuer_name(ca_cert.subject) + .public_key(client_key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - datetime.timedelta(minutes=5)) + .not_valid_after(now + datetime.timedelta(days=days_valid)) + .add_extension(x509.BasicConstraints(ca=False, path_length=None), critical=True) + .add_extension( + x509.KeyUsage( + digital_signature=True, + key_encipherment=True, + key_cert_sign=False, + crl_sign=False, + content_commitment=False, + data_encipherment=False, + key_agreement=False, + encipher_only=False, + decipher_only=False + ), + critical=True + ) + .add_extension( + x509.ExtendedKeyUsage([ExtendedKeyUsageOID.CLIENT_AUTH]), + critical=False + ) + .add_extension( + x509.SubjectKeyIdentifier.from_public_key(client_key.public_key()), + critical=False + ) + .add_extension( + x509.AuthorityKeyIdentifier.from_issuer_public_key(ca_key.public_key()), + critical=False + ) + .sign(ca_key, hashes.SHA256()) + ) + + cert_pem = cert.public_bytes(serialization.Encoding.PEM).decode("utf-8") + key_pem = client_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.TraditionalOpenSSL, + encryption_algorithm=serialization.NoEncryption() + ).decode("utf-8") + + return cert_pem, key_pem