use super::OpensslKeyError;
use crate::crypto::composite_kem::{
combiner_input, split_composite_kem_ciphertext, split_composite_kem_privkey,
split_composite_kem_spki_content, CompositeKemSpec, TradKemAlg,
};
use crate::crypto::composite_mldsa::{extract_spki_bitstring_payload, pkcs8_private_key_content};
use native_ossl::pkey::{DeriveCtx, KeygenCtx, Pkey, Private, Public};
use native_ossl::util::SecretBuf;
fn mlkem_name_cstr(variant: &str) -> Result<&'static std::ffi::CStr, OpensslKeyError> {
match variant {
"ML-KEM-768" => Ok(c"ML-KEM-768"),
"ML-KEM-1024" => Ok(c"ML-KEM-1024"),
other => Err(OpensslKeyError(format!(
"unsupported ML-KEM variant: {other}"
))),
}
}
fn mlkem_oid_for(variant: &str) -> Result<&'static [u32], OpensslKeyError> {
match variant {
"ML-KEM-768" => Ok(crate::oids::ML_KEM_768),
"ML-KEM-1024" => Ok(crate::oids::ML_KEM_1024),
other => Err(OpensslKeyError(format!(
"unsupported ML-KEM variant: {other}"
))),
}
}
fn ec_scalar_len(curve: &str) -> Result<usize, OpensslKeyError> {
match curve {
"P-256" => Ok(32),
"P-384" => Ok(48),
"P-521" => Ok(66),
"brainpoolP256r1" => Ok(32),
"brainpoolP384r1" => Ok(48),
other => Err(OpensslKeyError(format!("unsupported EC curve: {other}"))),
}
}
fn ec_composite_privkey_der(d: &[u8], curve: &str) -> Result<Vec<u8>, OpensslKeyError> {
use synta::tag::{Tag, TAG_SEQUENCE};
use synta::types::string::OctetStringRef;
use synta::{Encoder, Encoding, Integer, ObjectIdentifier};
let curve_oid_comps = super::composite::ec_curve_oid_comps(curve)?;
let curve_oid = ObjectIdentifier::new(curve_oid_comps)
.map_err(|e| OpensslKeyError(format!("invalid curve OID: {e}")))?;
let d_octet = OctetStringRef::new(d);
(|| -> synta::Result<Vec<u8>> {
let mut enc = Encoder::new(Encoding::Der);
enc.start_constructed_no_guard(Tag::universal_constructed(TAG_SEQUENCE))?;
enc.encode(&Integer::from_i64(1))?;
enc.encode(&d_octet)?;
enc.start_constructed_no_guard(Tag::context_specific_constructed(0))?;
enc.encode(&curve_oid)?;
enc.end_constructed()?;
enc.end_constructed()?;
enc.finish()
})()
.map_err(|e| OpensslKeyError(format!("composite ECPrivateKey DER encoding failed: {e}")))
}
fn trad_privkey_to_composite_sk(
trad_alg: &TradKemAlg,
pkey: &Pkey<Private>,
) -> Result<SecretBuf, OpensslKeyError> {
match trad_alg {
TradKemAlg::Rsa { .. } => Ok(SecretBuf::new(
pkcs8_private_key_content(&pkey.to_pkcs8_der()?).map_err(OpensslKeyError)?,
)),
TradKemAlg::X25519 | TradKemAlg::X448 => {
let exported = pkey.export()?;
let raw = exported
.get_octet_string(c"priv")
.map_err(|e| OpensslKeyError(format!("failed to export priv: {e}")))?;
Ok(SecretBuf::from_slice(raw))
}
TradKemAlg::Ec { curve } => {
let exported = pkey.export()?;
let scalar_len = ec_scalar_len(curve)?;
let raw = SecretBuf::new(
exported
.get_bn(c"priv")
.map_err(|e| OpensslKeyError(format!("failed to export EC priv: {e}")))?,
);
let raw_bytes = raw.as_ref();
if raw_bytes.len() > scalar_len {
return Err(OpensslKeyError(format!(
"EC private scalar longer than curve width ({} > {scalar_len})",
raw_bytes.len()
)));
}
let mut d = SecretBuf::with_len(scalar_len);
d.as_mut_slice()[scalar_len - raw_bytes.len()..].copy_from_slice(raw_bytes);
Ok(SecretBuf::new(ec_composite_privkey_der(d.as_ref(), curve)?))
}
}
}
fn trad_privkey_from_composite_sk(
trad_alg: &TradKemAlg,
trad_sk: &[u8],
) -> Result<Pkey<Private>, OpensslKeyError> {
match trad_alg {
TradKemAlg::Rsa { .. } => {
use synta::{Element, Null};
let pkcs8 = super::composite::encode_standalone_pkcs8(
crate::oids::RSA_ENCRYPTION,
Some(Element::Null(Null)),
trad_sk,
)?;
Ok(Pkey::<Private>::from_der(&pkcs8)?)
}
TradKemAlg::X25519 => {
use native_ossl::params::ParamBuilder;
use native_ossl::typed_params::curve25519;
let params = ParamBuilder::new()?
.set(curve25519::PRIV_KEY, trad_sk)?
.build()?;
Ok(Pkey::<Private>::from_params(None, c"X25519", ¶ms)?)
}
TradKemAlg::X448 => {
use native_ossl::params::ParamBuilder;
use native_ossl::typed_params::curve25519;
let params = ParamBuilder::new()?
.set(curve25519::PRIV_KEY, trad_sk)?
.build()?;
Ok(Pkey::<Private>::from_params(None, c"X448", ¶ms)?)
}
TradKemAlg::Ec { curve } => {
let params = super::composite::ec_alg_params(curve)?;
let pkcs8 = super::composite::encode_standalone_pkcs8(
crate::oids::EC_PUBLIC_KEY,
Some(params),
trad_sk,
)?;
Ok(Pkey::<Private>::from_der(&pkcs8)?)
}
}
}
fn encode_trad_kem_spki(trad_alg: &TradKemAlg, raw_pk: &[u8]) -> Result<Vec<u8>, OpensslKeyError> {
match trad_alg {
TradKemAlg::Rsa { .. } => {
use synta::{Element, Null};
super::composite::encode_standalone_spki(
crate::oids::RSA_ENCRYPTION,
Some(Element::Null(Null)),
raw_pk,
)
}
TradKemAlg::Ec { curve } => {
let params = super::composite::ec_alg_params(curve)?;
super::composite::encode_standalone_spki(
crate::oids::EC_PUBLIC_KEY,
Some(params),
raw_pk,
)
}
TradKemAlg::X25519 => {
super::composite::encode_standalone_spki(crate::oids::X25519, None, raw_pk)
}
TradKemAlg::X448 => {
super::composite::encode_standalone_spki(crate::oids::X448, None, raw_pk)
}
}
}
fn kem_combine_hash(input: &[u8]) -> Result<Vec<u8>, OpensslKeyError> {
let md = native_ossl::digest::DigestAlg::fetch(c"SHA3-256", None)
.map_err(|e| OpensslKeyError(format!("SHA3-256 not available: {e}")))?;
Ok(md.digest_to_vec(input)?)
}
fn generate_trad_kem_pkey(trad_alg: &TradKemAlg) -> Result<Pkey<Private>, OpensslKeyError> {
use native_ossl::params::ParamBuilder;
use native_ossl::typed_params::{ec, keygen};
match trad_alg {
TradKemAlg::Rsa { bits } => {
let params = ParamBuilder::new()?
.set(keygen::BITS, bits)?
.push_uint(c"e", 65537u32)?
.build()?;
let mut kgen = KeygenCtx::new(c"RSA")?;
kgen.set_params(¶ms)?;
Ok(kgen.generate()?)
}
TradKemAlg::Ec { curve } => {
let curve_cstr: &std::ffi::CStr = match *curve {
"P-256" => c"P-256",
"P-384" => c"P-384",
"P-521" => c"P-521",
"brainpoolP256r1" => c"brainpoolP256r1",
"brainpoolP384r1" => c"brainpoolP384r1",
other => return Err(OpensslKeyError(format!("unsupported EC curve: {other}"))),
};
let params = ParamBuilder::new()?.set(ec::GROUP, curve_cstr)?.build()?;
let mut kgen = KeygenCtx::new(c"EC")?;
kgen.set_params(¶ms)?;
Ok(kgen.generate()?)
}
TradKemAlg::X25519 => {
let mut kgen = KeygenCtx::new(c"X25519")?;
Ok(kgen.generate()?)
}
TradKemAlg::X448 => {
let mut kgen = KeygenCtx::new(c"X448")?;
Ok(kgen.generate()?)
}
}
}
fn generate_trad_kem_key(trad_alg: &TradKemAlg) -> Result<(SecretBuf, Vec<u8>), OpensslKeyError> {
let pkey = generate_trad_kem_pkey(trad_alg)?;
let spki_der = pkey.public_key_to_der()?;
let trad_sk = trad_privkey_to_composite_sk(trad_alg, &pkey)?;
let trad_pk = extract_spki_bitstring_payload(&spki_der).map_err(OpensslKeyError)?;
Ok((trad_sk, trad_pk))
}
fn trad_kem_encapsulate(
trad_alg: &TradKemAlg,
trad_pk: &[u8],
) -> Result<(Vec<u8>, SecretBuf), OpensslKeyError> {
match trad_alg {
TradKemAlg::Rsa { .. } => {
let spki = encode_trad_kem_spki(trad_alg, trad_pk)?;
let pub_pkey = Pkey::<Public>::from_der(&spki)?;
let mut shared_secret = SecretBuf::with_len(32);
native_ossl::rand::Rand::fill_private(shared_secret.as_mut_slice())?;
let ct = super::key_transport::rsa_oaep_encrypt_with_key(
&pub_pkey,
shared_secret.as_ref(),
"sha256",
)?;
Ok((ct, shared_secret))
}
TradKemAlg::Ec { .. } | TradKemAlg::X25519 | TradKemAlg::X448 => {
let trad_static_spki = encode_trad_kem_spki(trad_alg, trad_pk)?;
let trad_static_pkey = Pkey::<Public>::from_der(&trad_static_spki)?;
let eph_pkey = generate_trad_kem_pkey(trad_alg)?;
let trad_ct = extract_spki_bitstring_payload(&eph_pkey.public_key_to_der()?)
.map_err(OpensslKeyError)?;
let mut derive_ctx = DeriveCtx::new(&eph_pkey)?;
derive_ctx.set_peer(&trad_static_pkey)?;
let len = derive_ctx.derive_len()?;
let mut trad_ss = SecretBuf::with_len(len);
derive_ctx.derive(trad_ss.as_mut_slice())?;
Ok((trad_ct, trad_ss))
}
}
}
fn trad_kem_decapsulate(
trad_alg: &TradKemAlg,
trad_pkey: &Pkey<Private>,
trad_ct: &[u8],
) -> Result<SecretBuf, OpensslKeyError> {
match trad_alg {
TradKemAlg::Rsa { .. } => Ok(SecretBuf::new(
super::key_transport::rsa_oaep_decrypt_with_key(trad_pkey, trad_ct, "sha256")?,
)),
TradKemAlg::Ec { .. } | TradKemAlg::X25519 | TradKemAlg::X448 => {
let eph_spki = encode_trad_kem_spki(trad_alg, trad_ct)?;
let eph_pkey = Pkey::<Public>::from_der(&eph_spki)?;
let mut derive_ctx = DeriveCtx::new(trad_pkey)?;
derive_ctx.set_peer(&eph_pkey)?;
let len = derive_ctx.derive_len()?;
let mut trad_ss = SecretBuf::with_len(len);
derive_ctx.derive(trad_ss.as_mut_slice())?;
Ok(trad_ss)
}
}
}
pub(crate) fn priv_generate_composite_kem(
sub_arc: u32,
) -> Result<crate::crypto::BackendPrivateKey, OpensslKeyError> {
let spec = crate::crypto::composite_kem::composite_spec(sub_arc)
.ok_or_else(|| OpensslKeyError(format!("unknown composite ML-KEM sub-arc: {sub_arc}")))?;
let mlkem_name = mlkem_name_cstr(spec.mlkem_variant)?;
let mut kgen = KeygenCtx::new(mlkem_name)?;
let mlkem_pkey: Pkey<Private> = kgen.generate()?;
let mlkem_seed = SecretBuf::new({
let params = mlkem_pkey.export()?;
params
.get_octet_string(c"seed")
.map_err(|e| {
OpensslKeyError(format!("failed to export {} seed: {e}", spec.mlkem_variant))
})?
.to_vec()
});
let mlkem_spki = mlkem_pkey.public_key_to_der()?;
let mlkem_pk = extract_spki_bitstring_payload(&mlkem_spki).map_err(OpensslKeyError)?;
let (trad_sk, trad_pk) = generate_trad_kem_key(&spec.trad_alg)?;
let oid_comps = crate::crypto::composite_mldsa::composite_oid_components(spec.sub_arc);
let spki_der =
crate::crypto::composite_mldsa::encode_composite_spki(&oid_comps, &mlkem_pk, &trad_pk)
.map_err(OpensslKeyError)?;
let pkcs8_der = crate::crypto::composite_mldsa::encode_composite_pkcs8(
&oid_comps,
mlkem_seed.as_ref(),
trad_sk.as_ref(),
)
.map_err(OpensslKeyError)?;
let pkcs8_cell = std::sync::OnceLock::new();
pkcs8_cell.set(pkcs8_der).expect("fresh OnceLock");
Ok(crate::crypto::BackendPrivateKey {
pkcs8_der: pkcs8_cell,
spki_cache: Some(spki_der),
pkey: None,
pkcs11: None,
})
}
pub(crate) fn composite_kem_encapsulate_from_spki(
spki_der: &[u8],
spec: &'static CompositeKemSpec,
) -> Result<(Vec<u8>, Vec<u8>), OpensslKeyError> {
let payload = extract_spki_bitstring_payload(spki_der).map_err(OpensslKeyError)?;
let (mlkem_pk, trad_pk) =
split_composite_kem_spki_content(&payload, spec).map_err(OpensslKeyError)?;
let mlkem_oid = mlkem_oid_for(spec.mlkem_variant)?;
let mlkem_spki = super::composite::encode_standalone_spki(mlkem_oid, None, mlkem_pk)?;
let (mlkem_ct, mlkem_ss) = super::symmetric::pub_ml_kem_encapsulate(&mlkem_spki)?;
let mlkem_ss = SecretBuf::new(mlkem_ss);
let (trad_ct, trad_ss) = trad_kem_encapsulate(&spec.trad_alg, trad_pk)?;
let combined = SecretBuf::new(combiner_input(
mlkem_ss.as_ref(),
trad_ss.as_ref(),
&trad_ct,
trad_pk,
spec.label,
));
let ss = kem_combine_hash(combined.as_ref())?;
let mut ct = Vec::with_capacity(mlkem_ct.len() + trad_ct.len());
ct.extend_from_slice(&mlkem_ct);
ct.extend_from_slice(&trad_ct);
Ok((ct, ss))
}
pub(crate) fn composite_kem_decapsulate_from_pkcs8(
pkcs8_der: &[u8],
ciphertext: &[u8],
spec: &'static CompositeKemSpec,
) -> Result<Vec<u8>, OpensslKeyError> {
use native_ossl::params::ParamBuilder;
use native_ossl::typed_params::ml_kem;
let privkey_content =
SecretBuf::new(pkcs8_private_key_content(pkcs8_der).map_err(OpensslKeyError)?);
let (mlkem_seed, trad_sk) =
split_composite_kem_privkey(privkey_content.as_ref()).map_err(OpensslKeyError)?;
let (mlkem_ct, trad_ct) =
split_composite_kem_ciphertext(ciphertext, spec).map_err(OpensslKeyError)?;
let mlkem_name = mlkem_name_cstr(spec.mlkem_variant)?;
let seed_params = ParamBuilder::new()?
.set(ml_kem::SEED, mlkem_seed)?
.build()?;
let mlkem_pkey = Pkey::<Private>::from_params(None, mlkem_name, &seed_params)?;
let mlkem_pkcs8 = mlkem_pkey.to_pkcs8_der()?;
let mlkem_ss = SecretBuf::new(super::symmetric::priv_ml_kem_decapsulate(
&mlkem_pkcs8,
mlkem_ct,
)?);
let trad_pkey = trad_privkey_from_composite_sk(&spec.trad_alg, trad_sk)?;
let trad_pk =
extract_spki_bitstring_payload(&trad_pkey.public_key_to_der()?).map_err(OpensslKeyError)?;
let trad_ss = trad_kem_decapsulate(&spec.trad_alg, &trad_pkey, trad_ct)?;
let combined = SecretBuf::new(combiner_input(
mlkem_ss.as_ref(),
trad_ss.as_ref(),
trad_ct,
&trad_pk,
spec.label,
));
kem_combine_hash(combined.as_ref())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn composite_privkey_size_matches_spec() {
use crate::crypto::composite_kem::{
SUB_ARC_MLKEM1024_ECDH_BRAINPOOL_P384R1, SUB_ARC_MLKEM1024_ECDH_P384,
SUB_ARC_MLKEM1024_ECDH_P521, SUB_ARC_MLKEM1024_X448,
SUB_ARC_MLKEM768_ECDH_BRAINPOOL_P256R1, SUB_ARC_MLKEM768_ECDH_P256,
SUB_ARC_MLKEM768_ECDH_P384, SUB_ARC_MLKEM768_X25519,
};
for (sub_arc, expected_len) in [
(SUB_ARC_MLKEM768_X25519, 96usize),
(SUB_ARC_MLKEM768_ECDH_P256, 115),
(SUB_ARC_MLKEM768_ECDH_P384, 128),
(SUB_ARC_MLKEM768_ECDH_BRAINPOOL_P256R1, 116),
(SUB_ARC_MLKEM1024_ECDH_P384, 128),
(SUB_ARC_MLKEM1024_ECDH_BRAINPOOL_P384R1, 132),
(SUB_ARC_MLKEM1024_X448, 120),
(SUB_ARC_MLKEM1024_ECDH_P521, 146),
] {
let key = priv_generate_composite_kem(sub_arc).expect("keygen");
let der = key.pkcs8_der.get().expect("pkcs8 cached").clone();
let content = pkcs8_private_key_content(&der).expect("content");
assert_eq!(
content.len(),
expected_len,
"sub_arc {sub_arc}: composite private key size doesn't match Appendix A Table 3"
);
}
}
}