use alloc::vec::Vec;
use pq_oid::Algorithm;
use crate::asn1::length::{encode_length, encoded_length_size};
use crate::asn1::{algorithm, decode, encode, tags};
use crate::error::{Error, Result};
pub(crate) fn encode_pkcs8(algorithm: Algorithm, key_bytes: &[u8], out: &mut Vec<u8>) {
let version_len: usize = 3;
let alg_id_len = algorithm::encoded_algorithm_identifier_size(algorithm);
let octet_len = 1 + encoded_length_size(key_bytes.len()) + key_bytes.len();
let seq_content_len = version_len + alg_id_len + octet_len;
let total = 1 + encoded_length_size(seq_content_len) + seq_content_len;
out.reserve(total);
out.push(tags::TAG_SEQUENCE);
encode_length(seq_content_len, out);
encode::encode_integer_zero(out);
algorithm::encode_algorithm_identifier(algorithm, out);
encode::encode_octet_string(key_bytes, out);
}
pub(crate) fn decode_pkcs8(der: &[u8]) -> Result<(Algorithm, &[u8])> {
let outer = decode::read_tlv(der, 0)?;
if outer.tag != tags::TAG_SEQUENCE {
return Err(Error::InvalidDer("expected outer SEQUENCE in PKCS8"));
}
if outer.bytes_read != der.len() {
return Err(Error::InvalidDer(
"trailing data after outer SEQUENCE in PKCS8",
));
}
let seq = outer.value;
let ver_tlv = decode::read_tlv(seq, 0)?;
if ver_tlv.tag != tags::TAG_INTEGER {
return Err(Error::InvalidDer("expected INTEGER version in PKCS8"));
}
if ver_tlv.value.len() != 1 || (ver_tlv.value[0] != 0 && ver_tlv.value[0] != 1) {
return Err(Error::InvalidDer(
"unsupported PrivateKeyInfo version in PKCS8",
));
}
let version = ver_tlv.value[0];
let (alg, alg_bytes_read) = algorithm::decode_algorithm_identifier(seq, ver_tlv.bytes_read)?;
let key_offset = ver_tlv.bytes_read + alg_bytes_read;
let key_tlv = decode::read_tlv(seq, key_offset)?;
if key_tlv.tag != tags::TAG_OCTET_STRING {
return Err(Error::InvalidDer("expected OCTET STRING in PKCS8"));
}
let trailing_offset = key_offset + key_tlv.bytes_read;
let mut offset = trailing_offset;
while offset < seq.len() {
let trailing_tlv = decode::read_tlv(seq, offset)?;
if trailing_tlv.tag == tags::TAG_CONTEXT_0 {
} else if trailing_tlv.tag == tags::TAG_CONTEXT_1
|| trailing_tlv.tag == tags::TAG_CONTEXT_1_IMPLICIT
{
if version == 0 {
return Err(Error::InvalidDer(
"[1] publicKey not allowed in version 0 PKCS8",
));
}
} else {
return Err(Error::InvalidDer(
"unexpected trailing data in PKCS8 SEQUENCE",
));
}
offset += trailing_tlv.bytes_read;
}
let normalized = normalize_private_key_bytes(alg, key_tlv.value);
Ok((alg, normalized))
}
pub(crate) fn normalize_private_key_bytes(algorithm: Algorithm, bytes: &[u8]) -> &[u8] {
let expected = algorithm.private_key_size();
if bytes.len() == expected {
return bytes;
}
let tlv = match decode::read_tlv(bytes, 0) {
Ok(tlv) => tlv,
Err(_) => return bytes,
};
if tlv.tag == tags::TAG_OCTET_STRING
&& tlv.bytes_read == bytes.len()
&& tlv.value.len() == expected
{
return tlv.value;
}
if tlv.tag == tags::TAG_SEQUENCE && tlv.bytes_read == bytes.len() {
let seq_value = tlv.value;
let mut scan_offset = 0;
while scan_offset < seq_value.len() {
let inner = match decode::read_tlv(seq_value, scan_offset) {
Ok(inner) => inner,
Err(_) => break,
};
if inner.tag == tags::TAG_OCTET_STRING && inner.value.len() == expected {
return inner.value;
}
scan_offset += inner.bytes_read;
}
}
bytes
}
#[cfg(test)]
mod tests {
use super::*;
use crate::asn1::length::encode_length;
use pq_oid::{MlDsa, MlKem, SlhDsa};
#[test]
fn test_roundtrip_ml_kem_512() {
let alg = Algorithm::MlKem(MlKem::Kem512);
let key_bytes = vec![0xABu8; 1632];
let mut buf = Vec::new();
encode_pkcs8(alg, &key_bytes, &mut buf);
let (decoded_alg, decoded_bytes) = decode_pkcs8(&buf).unwrap();
assert_eq!(decoded_alg, alg);
assert_eq!(decoded_bytes, &key_bytes[..]);
}
#[test]
fn test_roundtrip_ml_dsa_44() {
let alg = Algorithm::MlDsa(MlDsa::Dsa44);
let key_bytes = vec![0xCDu8; 2560];
let mut buf = Vec::new();
encode_pkcs8(alg, &key_bytes, &mut buf);
let (decoded_alg, decoded_bytes) = decode_pkcs8(&buf).unwrap();
assert_eq!(decoded_alg, alg);
assert_eq!(decoded_bytes, &key_bytes[..]);
}
#[test]
fn test_roundtrip_slh_dsa_sha2_128s() {
let alg = Algorithm::SlhDsa(SlhDsa::Sha2_128s);
let key_bytes = vec![0xEFu8; 64];
let mut buf = Vec::new();
encode_pkcs8(alg, &key_bytes, &mut buf);
let (decoded_alg, decoded_bytes) = decode_pkcs8(&buf).unwrap();
assert_eq!(decoded_alg, alg);
assert_eq!(decoded_bytes, &key_bytes[..]);
}
#[test]
fn test_roundtrip_all_algorithms() {
for alg in Algorithm::all() {
let key_bytes = vec![0x42u8; alg.private_key_size()];
let mut buf = Vec::new();
encode_pkcs8(alg, &key_bytes, &mut buf);
let (decoded_alg, decoded_bytes) = decode_pkcs8(&buf).unwrap();
assert_eq!(decoded_alg, alg, "failed for {}", alg);
assert_eq!(decoded_bytes.len(), key_bytes.len());
}
}
#[test]
fn test_normalize_direct_raw() {
let alg = Algorithm::MlKem(MlKem::Kem512);
let raw = vec![0xAAu8; 1632];
let result = normalize_private_key_bytes(alg, &raw);
assert_eq!(result, &raw[..]);
}
#[test]
fn test_normalize_rfc8410_octet_string() {
let alg = Algorithm::SlhDsa(SlhDsa::Sha2_128s);
let raw = vec![0xBBu8; 64];
let mut wrapped = Vec::new();
wrapped.push(tags::TAG_OCTET_STRING);
encode_length(raw.len(), &mut wrapped);
wrapped.extend_from_slice(&raw);
let result = normalize_private_key_bytes(alg, &wrapped);
assert_eq!(result, &raw[..]);
}
#[test]
fn test_normalize_openssl_sequence() {
let alg = Algorithm::SlhDsa(SlhDsa::Sha2_128s);
let seed = vec![0xCCu8; 32]; let expanded = vec![0xDDu8; 64];
let mut inner = Vec::new();
inner.push(tags::TAG_OCTET_STRING);
encode_length(seed.len(), &mut inner);
inner.extend_from_slice(&seed);
inner.push(tags::TAG_OCTET_STRING);
encode_length(expanded.len(), &mut inner);
inner.extend_from_slice(&expanded);
let mut wrapped = Vec::new();
wrapped.push(tags::TAG_SEQUENCE);
encode_length(inner.len(), &mut wrapped);
wrapped.extend_from_slice(&inner);
let result = normalize_private_key_bytes(alg, &wrapped);
assert_eq!(result, &expanded[..]);
}
#[test]
fn test_normalize_fallback() {
let alg = Algorithm::MlKem(MlKem::Kem512);
let garbage = vec![0xFFu8; 100];
let result = normalize_private_key_bytes(alg, &garbage);
assert_eq!(result, &garbage[..]);
}
#[test]
fn test_version0_rejects_public_key_field() {
let alg = Algorithm::SlhDsa(SlhDsa::Sha2_128s);
let key_bytes = vec![0xAAu8; 64];
let mut valid = Vec::new();
encode_pkcs8(alg, &key_bytes, &mut valid);
let outer = decode::read_tlv(&valid, 0).unwrap();
let mut inner = outer.value.to_vec();
let pub_key = [0x01u8; 32];
inner.push(tags::TAG_CONTEXT_1);
encode_length(pub_key.len(), &mut inner);
inner.extend_from_slice(&pub_key);
let mut buf = Vec::new();
buf.push(tags::TAG_SEQUENCE);
encode_length(inner.len(), &mut buf);
buf.extend_from_slice(&inner);
let err = decode_pkcs8(&buf).unwrap_err();
assert!(
matches!(err, Error::InvalidDer(msg) if msg.contains("[1] publicKey")),
"expected [1] publicKey rejection, got: {:?}",
err
);
}
#[test]
fn test_version1_accepts_public_key_field() {
let alg = Algorithm::SlhDsa(SlhDsa::Sha2_128s);
let key_bytes = vec![0xBBu8; 64];
let mut inner = Vec::new();
inner.extend_from_slice(&[0x02, 0x01, 0x01]);
let mut alg_id = Vec::new();
algorithm::encode_algorithm_identifier(alg, &mut alg_id);
inner.extend_from_slice(&alg_id);
let mut octet = Vec::new();
encode::encode_octet_string(&key_bytes, &mut octet);
inner.extend_from_slice(&octet);
let pub_key = [0x01u8; 32];
inner.push(tags::TAG_CONTEXT_1);
encode_length(pub_key.len(), &mut inner);
inner.extend_from_slice(&pub_key);
let mut buf = Vec::new();
buf.push(tags::TAG_SEQUENCE);
encode_length(inner.len(), &mut buf);
buf.extend_from_slice(&inner);
let (decoded_alg, decoded_bytes) = decode_pkcs8(&buf).unwrap();
assert_eq!(decoded_alg, alg);
assert_eq!(decoded_bytes, &key_bytes[..]);
}
#[test]
fn test_version1_accepts_implicit_public_key_tag() {
let alg = Algorithm::SlhDsa(SlhDsa::Sha2_128s);
let key_bytes = vec![0xBBu8; 64];
let mut inner = Vec::new();
inner.extend_from_slice(&[0x02, 0x01, 0x01]);
let mut alg_id = Vec::new();
algorithm::encode_algorithm_identifier(alg, &mut alg_id);
inner.extend_from_slice(&alg_id);
let mut octet = Vec::new();
encode::encode_octet_string(&key_bytes, &mut octet);
inner.extend_from_slice(&octet);
let pub_key = [0x01u8; 32];
inner.push(tags::TAG_CONTEXT_1_IMPLICIT);
encode_length(pub_key.len(), &mut inner);
inner.extend_from_slice(&pub_key);
let mut buf = Vec::new();
buf.push(tags::TAG_SEQUENCE);
encode_length(inner.len(), &mut buf);
buf.extend_from_slice(&inner);
let (decoded_alg, decoded_bytes) = decode_pkcs8(&buf).unwrap();
assert_eq!(decoded_alg, alg);
assert_eq!(decoded_bytes, &key_bytes[..]);
}
#[test]
fn test_version0_rejects_implicit_public_key_tag() {
let alg = Algorithm::SlhDsa(SlhDsa::Sha2_128s);
let key_bytes = vec![0xAAu8; 64];
let mut valid = Vec::new();
encode_pkcs8(alg, &key_bytes, &mut valid);
let outer = decode::read_tlv(&valid, 0).unwrap();
let mut inner = outer.value.to_vec();
let pub_key = [0x01u8; 32];
inner.push(tags::TAG_CONTEXT_1_IMPLICIT);
encode_length(pub_key.len(), &mut inner);
inner.extend_from_slice(&pub_key);
let mut buf = Vec::new();
buf.push(tags::TAG_SEQUENCE);
encode_length(inner.len(), &mut buf);
buf.extend_from_slice(&inner);
let err = decode_pkcs8(&buf).unwrap_err();
assert!(
matches!(err, Error::InvalidDer(msg) if msg.contains("[1] publicKey")),
"expected [1] publicKey rejection, got: {:?}",
err
);
}
#[test]
fn test_version0_accepts_attributes_field() {
let alg = Algorithm::SlhDsa(SlhDsa::Sha2_128s);
let key_bytes = vec![0xCCu8; 64];
let mut valid = Vec::new();
encode_pkcs8(alg, &key_bytes, &mut valid);
let outer = decode::read_tlv(&valid, 0).unwrap();
let mut inner = outer.value.to_vec();
let attrs = [0x05, 0x00]; inner.push(tags::TAG_CONTEXT_0);
encode_length(attrs.len(), &mut inner);
inner.extend_from_slice(&attrs);
let mut buf = Vec::new();
buf.push(tags::TAG_SEQUENCE);
encode_length(inner.len(), &mut buf);
buf.extend_from_slice(&inner);
let (decoded_alg, decoded_bytes) = decode_pkcs8(&buf).unwrap();
assert_eq!(decoded_alg, alg);
assert_eq!(decoded_bytes, &key_bytes[..]);
}
#[test]
fn test_decode_invalid_version() {
let mut inner = Vec::new();
inner.extend_from_slice(&[0x02, 0x01, 0x02]);
let mut buf = Vec::new();
buf.push(tags::TAG_SEQUENCE);
encode_length(inner.len(), &mut buf);
buf.extend_from_slice(&inner);
assert!(decode_pkcs8(&buf).is_err());
}
}