use std::sync::Mutex;
use cryptoki::{
context::{CInitializeArgs, Pkcs11},
mechanism::Mechanism,
object::{Attribute, AttributeType, KeyType, ObjectClass, ObjectHandle},
session::{Session, UserType},
types::AuthPin,
};
use crate::{certificate::Certificate, digest::DigestAlgorithm, error::AdesError};
#[cfg(feature = "pkcs11")]
enum Pkcs11KeyType {
Rsa,
Ec,
}
#[cfg(feature = "pkcs11")]
pub struct Pkcs11Signer {
_pkcs11: Pkcs11,
session: Mutex<Session>,
key_handle: ObjectHandle,
certificate: Certificate,
digest: DigestAlgorithm,
key_type: Pkcs11KeyType,
}
#[cfg(feature = "pkcs11")]
impl Pkcs11Signer {
pub fn new(
lib_path: impl AsRef<std::path::Path>,
slot: u64,
pin: &str,
label: Option<&str>,
) -> Result<Self, AdesError> {
let pkcs11 = Pkcs11::new(lib_path).map_err(pkcs11_err)?;
pkcs11
.initialize(CInitializeArgs::OsThreads)
.map_err(pkcs11_err)?;
let slots = pkcs11.get_slots_with_token().map_err(pkcs11_err)?;
let slot = slots
.into_iter()
.nth(slot as usize)
.ok_or_else(|| AdesError::Pkcs11(format!("slot {slot} not found")))?;
let session = pkcs11.open_rw_session(slot).map_err(pkcs11_err)?;
let auth_pin = AuthPin::new(pin.to_owned());
session
.login(UserType::User, Some(&auth_pin))
.map_err(pkcs11_err)?;
let key_handle = find_object(&session, ObjectClass::PRIVATE_KEY, label)?;
let key_type = detect_key_type(&session, key_handle)?;
let cert_handle = find_object(&session, ObjectClass::CERTIFICATE, label)?;
let attrs = session
.get_attributes(cert_handle, &[AttributeType::Value])
.map_err(pkcs11_err)?;
let cert_der = attrs
.into_iter()
.find_map(|a| {
if let Attribute::Value(v) = a {
Some(v)
} else {
None
}
})
.ok_or_else(|| AdesError::Pkcs11("certificate object has no DER value".to_owned()))?;
let certificate = Certificate::from_der(&cert_der)?;
Ok(Self {
_pkcs11: pkcs11,
session: Mutex::new(session),
key_handle,
certificate,
digest: DigestAlgorithm::Sha256,
key_type,
})
}
pub fn list_slots(lib_path: impl AsRef<std::path::Path>) -> Result<Vec<u64>, AdesError> {
let pkcs11 = Pkcs11::new(lib_path).map_err(pkcs11_err)?;
pkcs11
.initialize(CInitializeArgs::OsThreads)
.map_err(pkcs11_err)?;
let slots = pkcs11.get_slots_with_token().map_err(pkcs11_err)?;
Ok(slots
.into_iter()
.enumerate()
.map(|(i, _)| i as u64)
.collect())
}
}
#[cfg(feature = "pkcs11")]
impl crate::signer::Signer for Pkcs11Signer {
type Error = AdesError;
fn sign_digest(&self, digest: &[u8]) -> Result<Vec<u8>, Self::Error> {
let session = self
.session
.lock()
.map_err(|_| AdesError::Pkcs11("session mutex poisoned".to_owned()))?;
match self.key_type {
Pkcs11KeyType::Rsa => {
let digest_info = build_digest_info(digest, self.digest)?;
session
.sign(&Mechanism::RsaPkcs, self.key_handle, &digest_info)
.map_err(pkcs11_err)
}
Pkcs11KeyType::Ec => {
let raw = session
.sign(&Mechanism::Ecdsa, self.key_handle, digest)
.map_err(pkcs11_err)?;
ec_raw_sig_to_der(&raw)
}
}
}
fn certificate(&self) -> &Certificate {
&self.certificate
}
fn digest_algorithm(&self) -> DigestAlgorithm {
self.digest
}
}
fn pkcs11_err(e: impl std::fmt::Display) -> AdesError {
AdesError::Pkcs11(e.to_string())
}
fn find_object(
session: &Session,
class: ObjectClass,
label: Option<&str>,
) -> Result<ObjectHandle, AdesError> {
let mut template = vec![Attribute::Class(class)];
if let Some(lbl) = label {
template.push(Attribute::Label(lbl.as_bytes().to_vec()));
}
session
.find_objects(&template)
.map_err(pkcs11_err)?
.into_iter()
.next()
.ok_or_else(|| {
AdesError::Pkcs11(format!(
"no {class:?} object found on token{}",
label
.map(|l| format!(" with label '{l}'"))
.unwrap_or_default()
))
})
}
fn detect_key_type(
session: &Session,
key_handle: ObjectHandle,
) -> Result<Pkcs11KeyType, AdesError> {
let attrs = session
.get_attributes(key_handle, &[AttributeType::KeyType])
.map_err(pkcs11_err)?;
let kt = attrs
.into_iter()
.find_map(|a| {
if let Attribute::KeyType(kt) = a {
Some(kt)
} else {
None
}
})
.ok_or_else(|| AdesError::Pkcs11("could not read CKA_KEY_TYPE attribute".to_owned()))?;
if kt == KeyType::RSA {
Ok(Pkcs11KeyType::Rsa)
} else if kt == KeyType::EC {
Ok(Pkcs11KeyType::Ec)
} else {
Err(AdesError::Pkcs11(format!(
"unsupported key type {kt:?} — only RSA and EC are supported"
)))
}
}
fn ec_raw_sig_to_der(raw: &[u8]) -> Result<Vec<u8>, AdesError> {
if !raw.len().is_multiple_of(2) {
return Err(AdesError::Pkcs11(format!(
"ECDSA raw signature length {} is not even",
raw.len()
)));
}
let coord = raw.len() / 2;
Ok(der_sequence(&[
&der_integer(&raw[..coord]),
&der_integer(&raw[coord..]),
]))
}
fn der_integer(bytes: &[u8]) -> Vec<u8> {
let trimmed: &[u8] = match bytes.iter().position(|&b| b != 0) {
Some(i) => &bytes[i..],
None => &[0],
};
let mut value = if trimmed[0] & 0x80 != 0 {
let mut v = vec![0x00];
v.extend_from_slice(trimmed);
v
} else {
trimmed.to_vec()
};
let mut out = vec![0x02]; der_push_length(&mut out, value.len());
out.append(&mut value);
out
}
fn der_sequence(items: &[&[u8]]) -> Vec<u8> {
let payload: Vec<u8> = items.iter().flat_map(|s| s.iter().copied()).collect();
let mut out = vec![0x30]; der_push_length(&mut out, payload.len());
out.extend_from_slice(&payload);
out
}
fn der_push_length(out: &mut Vec<u8>, len: usize) {
if len < 128 {
out.push(len as u8);
} else if len < 256 {
out.extend_from_slice(&[0x81, len as u8]);
} else {
out.extend_from_slice(&[0x82, (len >> 8) as u8, len as u8]);
}
}
fn build_digest_info(digest: &[u8], algo: DigestAlgorithm) -> Result<Vec<u8>, AdesError> {
let prefix: &[u8] = match algo {
DigestAlgorithm::Sha256 => &[
0x30, 0x31, 0x30, 0x0d, 0x06, 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02, 0x01, 0x05, 0x00, 0x04, 0x20, ],
DigestAlgorithm::Sha384 => &[
0x30, 0x41, 0x30, 0x0d, 0x06, 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02,
0x02, 0x05, 0x00, 0x04, 0x30, ],
DigestAlgorithm::Sha512 => &[
0x30, 0x51, 0x30, 0x0d, 0x06, 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02,
0x03, 0x05, 0x00, 0x04, 0x40, ],
};
let mut out = prefix.to_vec();
out.extend_from_slice(digest);
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn digest_info_sha256_length() {
let hash = [0u8; 32];
let di = build_digest_info(&hash, DigestAlgorithm::Sha256).unwrap();
assert_eq!(di.len(), 51);
assert_eq!(di[0], 0x30); assert_eq!(di[1], 0x31); }
#[test]
fn digest_info_sha384_length() {
let hash = [0u8; 48];
let di = build_digest_info(&hash, DigestAlgorithm::Sha384).unwrap();
assert_eq!(di.len(), 67);
assert_eq!(di[0], 0x30);
assert_eq!(di[1], 0x41); }
#[test]
fn digest_info_sha512_length() {
let hash = [0u8; 64];
let di = build_digest_info(&hash, DigestAlgorithm::Sha512).unwrap();
assert_eq!(di.len(), 83);
assert_eq!(di[0], 0x30);
assert_eq!(di[1], 0x51); }
#[test]
fn digest_info_sha256_oid_bytes() {
let hash = [0xabu8; 32];
let di = build_digest_info(&hash, DigestAlgorithm::Sha256).unwrap();
assert_eq!(
&di[4..15],
&[0x06, 0x09, 0x60, 0x86, 0x48, 0x01, 0x65, 0x03, 0x04, 0x02, 0x01]
);
assert_eq!(&di[di.len() - 32..], &[0xabu8; 32]);
}
}