use crate::{Result, SodiumError};
pub const BYTES: usize = libsodium_sys::crypto_scalarmult_BYTES as usize;
pub const SCALARBYTES: usize = libsodium_sys::crypto_scalarmult_SCALARBYTES as usize;
pub mod curve25519;
pub mod ed25519;
pub mod ristretto255;
pub fn scalarmult(secret_key: &[u8], public_key: &[u8]) -> Result<[u8; BYTES]> {
if secret_key.len() != SCALARBYTES {
return Err(SodiumError::InvalidInput(format!(
"secret key must be exactly {SCALARBYTES} bytes"
)));
}
if public_key.len() != BYTES {
return Err(SodiumError::InvalidInput(format!(
"public key must be exactly {BYTES} bytes"
)));
}
let mut shared_secret = [0u8; BYTES];
let result = unsafe {
libsodium_sys::crypto_scalarmult(
shared_secret.as_mut_ptr(),
secret_key.as_ptr(),
public_key.as_ptr(),
)
};
if result != 0 {
return Err(SodiumError::OperationError("scalarmult failed".into()));
}
Ok(shared_secret)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_scalarmult() {
let mut secret_key = vec![0u8; SCALARBYTES];
secret_key[0] = 1;
let mut public_key = vec![0u8; BYTES];
public_key[0] = 9;
let shared_secret = scalarmult(&secret_key, &public_key).unwrap();
assert_eq!(shared_secret.len(), BYTES);
}
#[test]
fn test_curve25519() {
let mut secret_key = vec![0u8; curve25519::SCALARBYTES];
secret_key[0] = 1;
let mut public_key = vec![0u8; curve25519::BYTES];
public_key[0] = 9;
let shared_secret = curve25519::scalarmult(&secret_key, &public_key).unwrap();
assert_eq!(shared_secret.len(), curve25519::BYTES);
}
#[test]
fn test_ed25519() {
let mut secret_key = vec![0u8; ed25519::SCALARBYTES];
secret_key.iter_mut().enumerate().for_each(|(i, byte)| {
*byte = i as u8;
});
let mut public_key = vec![0u8; ed25519::BYTES];
public_key.iter_mut().enumerate().for_each(|(i, byte)| {
*byte = (i + 100) as u8;
});
match ed25519::scalarmult(&secret_key, &public_key) {
Ok(shared_secret) => {
assert_eq!(shared_secret.len(), ed25519::BYTES);
}
Err(_) => {
}
}
}
#[test]
fn test_ristretto255() {
let mut secret_key = vec![0u8; ristretto255::SCALARBYTES];
secret_key.iter_mut().enumerate().for_each(|(i, byte)| {
*byte = i as u8;
});
let mut public_key = vec![0u8; ristretto255::BYTES];
public_key.iter_mut().enumerate().for_each(|(i, byte)| {
*byte = (i + 100) as u8;
});
match ristretto255::scalarmult(&secret_key, &public_key) {
Ok(shared_secret) => {
assert_eq!(shared_secret.len(), ristretto255::BYTES);
}
Err(_) => {
}
}
}
}