use ledger_secure_sdk_sys::*;
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum MlDsaParam {
MlDsa44,
MlDsa65,
#[cfg(feature = "mldsa_87")]
MlDsa87,
}
impl MlDsaParam {
const fn as_c(self) -> MLDSA_param_t {
match self {
MlDsaParam::MlDsa44 => MLDSA_44,
MlDsaParam::MlDsa65 => MLDSA_65,
#[cfg(feature = "mldsa_87")]
MlDsaParam::MlDsa87 => MLDSA_87,
}
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum MlDsaPrehash {
Sha256,
Sha512,
Sha3_256,
Sha3_512,
Shake128,
Shake256,
}
impl MlDsaPrehash {
const fn as_c(self) -> MLDSA_prehash_t {
match self {
MlDsaPrehash::Sha256 => MLDSA_PREHASH_SHA256,
MlDsaPrehash::Sha512 => MLDSA_PREHASH_SHA512,
MlDsaPrehash::Sha3_256 => MLDSA_PREHASH_SHA3_256,
MlDsaPrehash::Sha3_512 => MLDSA_PREHASH_SHA3_512,
MlDsaPrehash::Shake128 => MLDSA_PREHASH_SHAKE128,
MlDsaPrehash::Shake256 => MLDSA_PREHASH_SHAKE256,
}
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum MlDsaError {
InvalidParameter,
InvalidParameterValue,
InternalError,
}
impl From<u32> for MlDsaError {
fn from(code: u32) -> Self {
match code {
CX_INVALID_PARAMETER => MlDsaError::InvalidParameter,
CX_INVALID_PARAMETER_VALUE => MlDsaError::InvalidParameterValue,
_ => MlDsaError::InternalError,
}
}
}
pub const MLDSA44_PK_LEN: usize = MLDSA44_PUBLICKEYBYTES as usize;
pub const MLDSA44_SK_LEN: usize = MLDSA44_SECRETKEYBYTES as usize;
pub const MLDSA44_SIG_LEN: usize = MLDSA44_SIGBYTES as usize;
pub const MLDSA65_PK_LEN: usize = MLDSA65_PUBLICKEYBYTES as usize;
pub const MLDSA65_SK_LEN: usize = MLDSA65_SECRETKEYBYTES as usize;
pub const MLDSA65_SIG_LEN: usize = MLDSA65_SIGBYTES as usize;
#[cfg(feature = "mldsa_87")]
pub const MLDSA87_PK_LEN: usize = MLDSA87_PUBLICKEYBYTES as usize;
#[cfg(feature = "mldsa_87")]
pub const MLDSA87_SK_LEN: usize = MLDSA87_SECRETKEYBYTES as usize;
#[cfg(feature = "mldsa_87")]
pub const MLDSA87_SIG_LEN: usize = MLDSA87_SIGBYTES as usize;
pub const MAX_CTX_LEN: usize = 255;
impl MlDsaParam {
pub const fn pk_len(self) -> usize {
match self {
MlDsaParam::MlDsa44 => MLDSA44_PK_LEN,
MlDsaParam::MlDsa65 => MLDSA65_PK_LEN,
#[cfg(feature = "mldsa_87")]
MlDsaParam::MlDsa87 => MLDSA87_PK_LEN,
}
}
pub const fn sk_len(self) -> usize {
match self {
MlDsaParam::MlDsa44 => MLDSA44_SK_LEN,
MlDsaParam::MlDsa65 => MLDSA65_SK_LEN,
#[cfg(feature = "mldsa_87")]
MlDsaParam::MlDsa87 => MLDSA87_SK_LEN,
}
}
pub const fn sig_len(self) -> usize {
match self {
MlDsaParam::MlDsa44 => MLDSA44_SIG_LEN,
MlDsaParam::MlDsa65 => MLDSA65_SIG_LEN,
#[cfg(feature = "mldsa_87")]
MlDsaParam::MlDsa87 => MLDSA87_SIG_LEN,
}
}
}
pub fn keygen(pk: &mut [u8], sk: &mut [u8], param: MlDsaParam) -> Result<(), MlDsaError> {
let err = unsafe {
MLDSA_keygen(
pk.as_mut_ptr(),
pk.len(),
sk.as_mut_ptr(),
sk.len(),
param.as_c(),
)
};
if err != CX_OK {
Err(err.into())
} else {
Ok(())
}
}
pub fn sign(
sig: &mut [u8],
msg: &[u8],
ctx: &[u8],
sk: &[u8],
param: MlDsaParam,
) -> Result<usize, MlDsaError> {
if ctx.len() > MAX_CTX_LEN {
return Err(MlDsaError::InvalidParameterValue);
}
let mut sig_actual_len: usize = 0;
let ctx_ptr = if ctx.is_empty() {
core::ptr::null()
} else {
ctx.as_ptr()
};
let err = unsafe {
MLDSA_sign(
sig.as_mut_ptr(),
sig.len(),
&mut sig_actual_len,
msg.as_ptr(),
msg.len(),
ctx_ptr,
ctx.len(),
sk.as_ptr(),
sk.len(),
param.as_c(),
)
};
if err != CX_OK {
Err(err.into())
} else {
Ok(sig_actual_len)
}
}
pub fn verify(
sig: &[u8],
msg: &[u8],
ctx: &[u8],
pk: &[u8],
param: MlDsaParam,
) -> Result<(), MlDsaError> {
if ctx.len() > MAX_CTX_LEN {
return Err(MlDsaError::InvalidParameterValue);
}
let ctx_ptr = if ctx.is_empty() {
core::ptr::null()
} else {
ctx.as_ptr()
};
let err = unsafe {
MLDSA_verify(
sig.as_ptr(),
sig.len(),
msg.as_ptr(),
msg.len(),
ctx_ptr,
ctx.len(),
pk.as_ptr(),
pk.len(),
param.as_c(),
)
};
if err != CX_OK {
Err(err.into())
} else {
Ok(())
}
}
pub fn sign_prehash(
sig: &mut [u8],
ph: &[u8],
ctx: &[u8],
sk: &[u8],
prehash_alg: MlDsaPrehash,
param: MlDsaParam,
) -> Result<usize, MlDsaError> {
if ctx.len() > MAX_CTX_LEN {
return Err(MlDsaError::InvalidParameterValue);
}
let mut sig_actual_len: usize = 0;
let ctx_ptr = if ctx.is_empty() {
core::ptr::null()
} else {
ctx.as_ptr()
};
let err = unsafe {
MLDSA_sign_prehash(
sig.as_mut_ptr(),
sig.len(),
&mut sig_actual_len,
ph.as_ptr(),
ph.len(),
ctx_ptr,
ctx.len(),
sk.as_ptr(),
sk.len(),
prehash_alg.as_c(),
param.as_c(),
)
};
if err != CX_OK {
Err(err.into())
} else {
Ok(sig_actual_len)
}
}
pub fn verify_prehash(
sig: &[u8],
ph: &[u8],
ctx: &[u8],
pk: &[u8],
prehash_alg: MlDsaPrehash,
param: MlDsaParam,
) -> Result<(), MlDsaError> {
if ctx.len() > MAX_CTX_LEN {
return Err(MlDsaError::InvalidParameterValue);
}
let ctx_ptr = if ctx.is_empty() {
core::ptr::null()
} else {
ctx.as_ptr()
};
let err = unsafe {
MLDSA_verify_prehash(
sig.as_ptr(),
sig.len(),
ph.as_ptr(),
ph.len(),
ctx_ptr,
ctx.len(),
pk.as_ptr(),
pk.len(),
prehash_alg.as_c(),
param.as_c(),
)
};
if err != CX_OK {
Err(err.into())
} else {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::assert_eq_err as assert_eq;
use crate::testing::TestType;
use testmacro::test_item as test;
const TEST_MSG: &[u8] = b"Test message";
#[test]
fn test_mldsa44_sign_verify() {
let mut pk = [0u8; MLDSA44_PK_LEN];
let mut sk = [0u8; MLDSA44_SK_LEN];
keygen(&mut pk, &mut sk, MlDsaParam::MlDsa44).unwrap();
let mut sig = [0u8; MLDSA44_SIG_LEN];
let sig_len = sign(&mut sig, TEST_MSG, &[], &sk, MlDsaParam::MlDsa44).unwrap();
assert_eq!(sig_len, MLDSA44_SIG_LEN);
verify(&sig[..sig_len], TEST_MSG, &[], &pk, MlDsaParam::MlDsa44).unwrap();
}
#[test]
fn test_mldsa65_sign_verify() {
let mut pk = [0u8; MLDSA65_PK_LEN];
let mut sk = [0u8; MLDSA65_SK_LEN];
keygen(&mut pk, &mut sk, MlDsaParam::MlDsa65).unwrap();
let mut sig = [0u8; MLDSA65_SIG_LEN];
let sig_len = sign(&mut sig, TEST_MSG, &[], &sk, MlDsaParam::MlDsa65).unwrap();
assert_eq!(sig_len, MLDSA65_SIG_LEN);
verify(&sig[..sig_len], TEST_MSG, &[], &pk, MlDsaParam::MlDsa65).unwrap();
}
#[test]
fn test_mldsa44_sign_verify_with_context() {
let mut pk = [0u8; MLDSA44_PK_LEN];
let mut sk = [0u8; MLDSA44_SK_LEN];
keygen(&mut pk, &mut sk, MlDsaParam::MlDsa44).unwrap();
let ctx = b"test context";
let mut sig = [0u8; MLDSA44_SIG_LEN];
let sig_len = sign(&mut sig, TEST_MSG, ctx, &sk, MlDsaParam::MlDsa44).unwrap();
verify(&sig[..sig_len], TEST_MSG, ctx, &pk, MlDsaParam::MlDsa44).unwrap();
}
}