use foreign_types::ForeignType;
use openssl::ec::{EcGroup, EcKey};
use openssl::error::ErrorStack;
use openssl::nid::Nid;
use openssl::pkey::{Id, PKey, Private};
use openssl::rsa::Rsa;
use openssl::x509::{X509, X509Req};
unsafe extern "C" {
pub fn X509_sign(
x: *mut openssl_sys::X509,
pkey: *mut openssl_sys::EVP_PKEY,
md: *const openssl_sys::EVP_MD,
) -> ::std::os::raw::c_int;
pub fn X509_sign_ctx(
x: *mut openssl_sys::X509,
ctx: *mut openssl_sys::EVP_MD_CTX,
) -> ::std::os::raw::c_int;
}
unsafe extern "C" {
pub fn X509_REQ_sign(
req: *mut openssl_sys::X509_REQ,
pkey: *mut openssl_sys::EVP_PKEY,
md: *const openssl_sys::EVP_MD,
) -> ::std::os::raw::c_int;
pub fn X509_REQ_sign_ctx(
req: *mut openssl_sys::X509_REQ,
ctx: *mut openssl_sys::EVP_MD_CTX,
) -> ::std::os::raw::c_int;
}
pub(crate) fn sign_certificate_digestless(
cert: &X509,
pkey: &PKey<openssl::pkey::Private>,
) -> Result<(), String> {
if !is_digestless_key(pkey) {
return Err("sign_certificate_digestless called with non-digestless key".to_string());
}
let cert_ptr = cert.as_ptr();
let pkey_ptr = pkey.as_ptr();
if pkey.id() == Id::ED25519 {
let result = unsafe { X509_sign(cert_ptr, pkey_ptr, std::ptr::null()) };
return if result > 0 {
Ok(())
} else {
Err("Failed to sign certificate with Ed25519".to_string())
};
}
let ctx = MdCtx(unsafe { openssl_sys::EVP_MD_CTX_new() });
if ctx.0.is_null() {
return Err("EVP_MD_CTX_new returned NULL".to_string());
}
let init = unsafe {
openssl_sys::EVP_DigestSignInit(
ctx.0,
std::ptr::null_mut(),
std::ptr::null(),
std::ptr::null_mut(),
pkey_ptr,
)
};
if init <= 0 {
return Err("EVP_DigestSignInit failed for PQC key".to_string());
}
let result = unsafe { X509_sign_ctx(cert_ptr, ctx.0) };
if result > 0 {
Ok(())
} else {
Err("X509_sign_ctx failed for PQC key".to_string())
}
}
pub(crate) fn sign_x509_req_digestless(req: &X509Req, pkey: &PKey<Private>) -> Result<(), String> {
if !is_digestless_key(pkey) {
return Err("sign_x509_req_digestless called with non-digestless key".to_string());
}
let req_ptr = req.as_ptr();
let pkey_ptr = pkey.as_ptr();
if pkey.id() == Id::ED25519 {
let result = unsafe { X509_REQ_sign(req_ptr, pkey_ptr, std::ptr::null()) };
return if result > 0 {
Ok(())
} else {
Err("Failed to sign X509Req with Ed25519".to_string())
};
}
let ctx = MdCtx(unsafe { openssl_sys::EVP_MD_CTX_new() });
if ctx.0.is_null() {
return Err("EVP_MD_CTX_new returned NULL".to_string());
}
let init = unsafe {
openssl_sys::EVP_DigestSignInit(
ctx.0,
std::ptr::null_mut(),
std::ptr::null(),
std::ptr::null_mut(),
pkey_ptr,
)
};
if init <= 0 {
return Err("EVP_DigestSignInit failed for PQC key".to_string());
}
let result = unsafe { X509_REQ_sign_ctx(req_ptr, ctx.0) };
if result > 0 {
Ok(())
} else {
Err("X509_REQ_sign_ctx failed for PQC key".to_string())
}
}
#[cfg(feature = "pqc")]
#[allow(dead_code)]
const ML_KEM_OIDS: [&str; 3] = [
"2.16.840.1.101.3.4.4.1", "2.16.840.1.101.3.4.4.2", "2.16.840.1.101.3.4.4.3", ];
struct MdCtx(*mut openssl_sys::EVP_MD_CTX);
impl Drop for MdCtx {
fn drop(&mut self) {
if !self.0.is_null() {
unsafe { openssl_sys::EVP_MD_CTX_free(self.0) }
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum KeyType {
RSA2048,
RSA4096,
P224,
P256,
P384,
P521,
Ed25519,
#[cfg(feature = "pqc")]
MlDsa44,
#[cfg(feature = "pqc")]
MlDsa65,
#[cfg(feature = "pqc")]
MlDsa87,
#[cfg(feature = "pqc")]
SlhDsaSha2_128s,
#[cfg(feature = "pqc")]
SlhDsaSha2_192s,
#[cfg(feature = "pqc")]
SlhDsaSha2_256s,
#[cfg(feature = "pqc")]
MlKem512,
#[cfg(feature = "pqc")]
MlKem768,
#[cfg(feature = "pqc")]
MlKem1024,
}
pub(crate) fn select_key(key_type: &Option<KeyType>) -> Result<PKey<Private>, ErrorStack> {
match key_type {
Some(KeyType::P224) => {
let group = EcGroup::from_curve_name(Nid::SECP224R1)?;
let ec_key = EcKey::generate(&group)?;
PKey::from_ec_key(ec_key)
}
Some(KeyType::P256) => {
let group = EcGroup::from_curve_name(Nid::X9_62_PRIME256V1)?;
let ec_key = EcKey::generate(&group)?;
PKey::from_ec_key(ec_key)
}
Some(KeyType::P384) => {
let group = EcGroup::from_curve_name(Nid::SECP384R1)?;
let ec_key = EcKey::generate(&group)?;
PKey::from_ec_key(ec_key)
}
Some(KeyType::P521) => {
let group = EcGroup::from_curve_name(Nid::SECP521R1)?;
let ec_key = EcKey::generate(&group)?;
PKey::from_ec_key(ec_key)
}
Some(KeyType::Ed25519) => PKey::generate_ed25519(),
#[cfg(feature = "pqc")]
Some(KeyType::MlDsa44) => generate_pqc_key("ML-DSA-44"),
#[cfg(feature = "pqc")]
Some(KeyType::MlDsa65) => generate_pqc_key("ML-DSA-65"),
#[cfg(feature = "pqc")]
Some(KeyType::MlDsa87) => generate_pqc_key("ML-DSA-87"),
#[cfg(feature = "pqc")]
Some(KeyType::SlhDsaSha2_128s) => generate_pqc_key("SLH-DSA-SHA2-128s"),
#[cfg(feature = "pqc")]
Some(KeyType::SlhDsaSha2_192s) => generate_pqc_key("SLH-DSA-SHA2-192s"),
#[cfg(feature = "pqc")]
Some(KeyType::SlhDsaSha2_256s) => generate_pqc_key("SLH-DSA-SHA2-256s"),
#[cfg(feature = "pqc")]
Some(KeyType::MlKem512) => generate_pqc_key("ML-KEM-512"),
#[cfg(feature = "pqc")]
Some(KeyType::MlKem768) => generate_pqc_key("ML-KEM-768"),
#[cfg(feature = "pqc")]
Some(KeyType::MlKem1024) => generate_pqc_key("ML-KEM-1024"),
Some(KeyType::RSA4096) => {
let rsa = Rsa::generate(4096)?;
PKey::from_rsa(rsa)
}
_ => {
let rsa = Rsa::generate(2048)?;
PKey::from_rsa(rsa)
}
}
}
#[cfg(feature = "pqc")]
mod pqc {
use foreign_types::ForeignType;
use openssl::error::ErrorStack;
use openssl::pkey::{PKey, Private};
use std::ffi::CString;
unsafe extern "C" {
fn EVP_PKEY_CTX_new_from_name(
libctx: *mut std::ffi::c_void,
name: *const std::os::raw::c_char,
propquery: *const std::os::raw::c_char,
) -> *mut openssl_sys::EVP_PKEY_CTX;
fn EVP_PKEY_keygen_init(ctx: *mut openssl_sys::EVP_PKEY_CTX) -> std::os::raw::c_int;
fn EVP_PKEY_generate(
ctx: *mut openssl_sys::EVP_PKEY_CTX,
ppkey: *mut *mut openssl_sys::EVP_PKEY,
) -> std::os::raw::c_int;
fn EVP_PKEY_CTX_free(ctx: *mut openssl_sys::EVP_PKEY_CTX);
pub fn EVP_PKEY_is_a(
pkey: *mut openssl_sys::EVP_PKEY,
name: *const std::os::raw::c_char,
) -> std::os::raw::c_int;
}
struct PkeyCtx(*mut openssl_sys::EVP_PKEY_CTX);
impl Drop for PkeyCtx {
fn drop(&mut self) {
if !self.0.is_null() {
unsafe { EVP_PKEY_CTX_free(self.0) }
}
}
}
pub(crate) fn generate_pqc_key(alg_name: &str) -> Result<PKey<Private>, ErrorStack> {
let cname = CString::new(alg_name).expect("alg_name contains interior NUL");
let ctx_ptr = unsafe {
EVP_PKEY_CTX_new_from_name(std::ptr::null_mut(), cname.as_ptr(), std::ptr::null())
};
if ctx_ptr.is_null() {
return Err(ErrorStack::get());
}
let ctx = PkeyCtx(ctx_ptr);
if unsafe { EVP_PKEY_keygen_init(ctx.0) } <= 0 {
return Err(ErrorStack::get());
}
let mut pkey_ptr: *mut openssl_sys::EVP_PKEY = std::ptr::null_mut();
if unsafe { EVP_PKEY_generate(ctx.0, &mut pkey_ptr) } <= 0 {
return Err(ErrorStack::get());
}
if pkey_ptr.is_null() {
return Err(ErrorStack::get());
}
Ok(unsafe { PKey::<Private>::from_ptr(pkey_ptr) })
}
}
#[cfg(feature = "pqc")]
pub(crate) use pqc::generate_pqc_key;
#[cfg(feature = "pqc")]
pub(crate) fn is_pqc_pkey<T>(pkey: &PKey<T>) -> bool {
use std::ffi::CString;
use std::sync::OnceLock;
static NAMES: OnceLock<[CString; 6]> = OnceLock::new();
let names = NAMES.get_or_init(|| {
[
CString::new("ML-DSA-44").unwrap(),
CString::new("ML-DSA-65").unwrap(),
CString::new("ML-DSA-87").unwrap(),
CString::new("SLH-DSA-SHA2-128s").unwrap(),
CString::new("SLH-DSA-SHA2-192s").unwrap(),
CString::new("SLH-DSA-SHA2-256s").unwrap(),
]
});
use foreign_types::ForeignType;
let ptr = pkey.as_ptr();
names
.iter()
.any(|n| unsafe { pqc::EVP_PKEY_is_a(ptr, n.as_ptr()) } == 1)
}
#[cfg(feature = "pqc")]
pub(crate) fn is_mlkem_pkey<T>(pkey: &PKey<T>) -> bool {
use std::ffi::CString;
use std::sync::OnceLock;
static NAMES: OnceLock<[CString; 3]> = OnceLock::new();
let names = NAMES.get_or_init(|| {
[
CString::new("ML-KEM-512").unwrap(),
CString::new("ML-KEM-768").unwrap(),
CString::new("ML-KEM-1024").unwrap(),
]
});
use foreign_types::ForeignType;
let ptr = pkey.as_ptr();
names
.iter()
.any(|n| unsafe { pqc::EVP_PKEY_is_a(ptr, n.as_ptr()) } == 1)
}
pub(crate) fn is_digestless_key(pkey: &PKey<Private>) -> bool {
if pkey.id() == Id::ED25519 {
return true;
}
#[cfg(feature = "pqc")]
{
return is_pqc_pkey(pkey);
}
#[allow(unreachable_code)]
false
}
#[cfg(feature = "pqc")]
pub(crate) fn reject_mlkem_signing(
pkey: &PKey<Private>,
message: &'static str,
) -> Result<(), Box<dyn std::error::Error>> {
if is_mlkem_pkey(pkey) {
return Err(message.into());
}
Ok(())
}