use crate::{Result, SodiumError};
use libsodium_sys;
use std::ffi::CStr;
pub const PUBLICKEYBYTES: usize = libsodium_sys::crypto_kx_PUBLICKEYBYTES as usize;
pub const SECRETKEYBYTES: usize = libsodium_sys::crypto_kx_SECRETKEYBYTES as usize;
pub const SEEDBYTES: usize = libsodium_sys::crypto_kx_SEEDBYTES as usize;
pub const SESSIONKEYBYTES: usize = libsodium_sys::crypto_kx_SESSIONKEYBYTES as usize;
pub const PRIMITIVE: &str = "x25519blake2b";
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct PublicKey([u8; PUBLICKEYBYTES]);
#[derive(Debug, Clone, Eq, PartialEq, zeroize::Zeroize, zeroize::ZeroizeOnDrop)]
pub struct SecretKey([u8; SECRETKEYBYTES]);
pub struct KeyPair {
pub public_key: PublicKey,
pub secret_key: SecretKey,
}
#[derive(Debug, Clone, Eq, PartialEq, zeroize::Zeroize, zeroize::ZeroizeOnDrop)]
pub struct SessionKeys {
pub tx: [u8; SESSIONKEYBYTES],
pub rx: [u8; SESSIONKEYBYTES],
}
impl PublicKey {
pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
if bytes.len() != PUBLICKEYBYTES {
return Err(SodiumError::InvalidInput(format!(
"public key must be exactly {PUBLICKEYBYTES} bytes"
)));
}
let mut key = [0u8; PUBLICKEYBYTES];
key.copy_from_slice(bytes);
Ok(PublicKey(key))
}
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
}
impl AsRef<[u8]> for PublicKey {
fn as_ref(&self) -> &[u8] {
self.as_bytes()
}
}
impl TryFrom<&[u8]> for PublicKey {
type Error = SodiumError;
fn try_from(bytes: &[u8]) -> std::result::Result<Self, Self::Error> {
Self::from_bytes(bytes)
}
}
impl From<[u8; PUBLICKEYBYTES]> for PublicKey {
fn from(bytes: [u8; PUBLICKEYBYTES]) -> Self {
PublicKey(bytes)
}
}
impl From<PublicKey> for [u8; PUBLICKEYBYTES] {
fn from(key: PublicKey) -> Self {
key.0
}
}
impl SecretKey {
pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
if bytes.len() != SECRETKEYBYTES {
return Err(SodiumError::InvalidInput(format!(
"secret key must be exactly {SECRETKEYBYTES} bytes"
)));
}
let mut key = [0u8; SECRETKEYBYTES];
key.copy_from_slice(bytes);
Ok(SecretKey(key))
}
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
}
impl AsRef<[u8]> for SecretKey {
fn as_ref(&self) -> &[u8] {
self.as_bytes()
}
}
impl TryFrom<&[u8]> for SecretKey {
type Error = SodiumError;
fn try_from(bytes: &[u8]) -> std::result::Result<Self, Self::Error> {
Self::from_bytes(bytes)
}
}
impl From<[u8; SECRETKEYBYTES]> for SecretKey {
fn from(bytes: [u8; SECRETKEYBYTES]) -> Self {
SecretKey(bytes)
}
}
impl From<SecretKey> for [u8; SECRETKEYBYTES] {
fn from(key: SecretKey) -> Self {
key.0
}
}
impl KeyPair {
pub fn generate() -> Result<Self> {
let mut pk = [0u8; PUBLICKEYBYTES];
let mut sk = [0u8; SECRETKEYBYTES];
let result = unsafe { libsodium_sys::crypto_kx_keypair(pk.as_mut_ptr(), sk.as_mut_ptr()) };
if result != 0 {
return Err(SodiumError::OperationError(
"failed to generate keypair".into(),
));
}
Ok(Self {
public_key: PublicKey(pk),
secret_key: SecretKey(sk),
})
}
pub fn from_seed(seed: &[u8]) -> Result<Self> {
if seed.len() != SEEDBYTES {
return Err(SodiumError::InvalidInput(format!(
"invalid seed length: expected {}, got {}",
SEEDBYTES,
seed.len()
)));
}
let mut pk = [0u8; PUBLICKEYBYTES];
let mut sk = [0u8; SECRETKEYBYTES];
let result = unsafe {
libsodium_sys::crypto_kx_seed_keypair(pk.as_mut_ptr(), sk.as_mut_ptr(), seed.as_ptr())
};
if result != 0 {
return Err(SodiumError::OperationError(
"failed to generate keypair from seed".into(),
));
}
Ok(Self {
public_key: PublicKey(pk),
secret_key: SecretKey(sk),
})
}
pub fn into_tuple(self) -> (PublicKey, SecretKey) {
(self.public_key, self.secret_key)
}
}
pub fn seedbytes() -> usize {
unsafe { libsodium_sys::crypto_kx_seedbytes() }
}
pub fn sessionkeybytes() -> usize {
unsafe { libsodium_sys::crypto_kx_sessionkeybytes() }
}
pub fn primitive() -> &'static str {
unsafe {
CStr::from_ptr(libsodium_sys::crypto_kx_primitive())
.to_str()
.expect("crypto_kx primitive should be valid UTF-8")
}
}
pub fn client_session_keys(
client_pk: &PublicKey,
client_sk: &SecretKey,
server_pk: &PublicKey,
) -> Result<SessionKeys> {
let mut rx = [0u8; SESSIONKEYBYTES];
let mut tx = [0u8; SESSIONKEYBYTES];
let result = unsafe {
libsodium_sys::crypto_kx_client_session_keys(
rx.as_mut_ptr(),
tx.as_mut_ptr(),
client_pk.as_bytes().as_ptr(),
client_sk.as_bytes().as_ptr(),
server_pk.as_bytes().as_ptr(),
)
};
if result != 0 {
return Err(SodiumError::OperationError(
"failed to compute session keys".into(),
));
}
Ok(SessionKeys { rx, tx })
}
pub fn server_session_keys(
server_pk: &PublicKey,
server_sk: &SecretKey,
client_pk: &PublicKey,
) -> Result<SessionKeys> {
let mut rx = [0u8; SESSIONKEYBYTES];
let mut tx = [0u8; SESSIONKEYBYTES];
let result = unsafe {
libsodium_sys::crypto_kx_server_session_keys(
rx.as_mut_ptr(),
tx.as_mut_ptr(),
server_pk.as_bytes().as_ptr(),
server_sk.as_bytes().as_ptr(),
client_pk.as_bytes().as_ptr(),
)
};
if result != 0 {
return Err(SodiumError::OperationError(
"failed to compute session keys".into(),
));
}
Ok(SessionKeys { rx, tx })
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_keypair_generation() {
let keypair = KeyPair::generate().unwrap();
let pk = keypair.public_key;
let sk = keypair.secret_key;
assert_eq!(pk.as_bytes().len(), PUBLICKEYBYTES);
assert_eq!(sk.as_bytes().len(), SECRETKEYBYTES);
}
#[test]
fn test_constants_and_primitive() {
assert_eq!(seedbytes(), SEEDBYTES);
assert_eq!(sessionkeybytes(), SESSIONKEYBYTES);
assert_eq!(primitive(), PRIMITIVE);
}
#[test]
fn test_seed_keypair() {
let seed = [42u8; SEEDBYTES];
let keypair1 = KeyPair::from_seed(&seed).unwrap();
let keypair2 = KeyPair::from_seed(&seed).unwrap();
assert_eq!(keypair1.public_key, keypair2.public_key);
assert_eq!(keypair1.secret_key, keypair2.secret_key);
}
#[test]
fn test_key_exchange() {
let client_keypair = KeyPair::generate().unwrap();
let client_pk = client_keypair.public_key;
let client_sk = client_keypair.secret_key;
let server_keypair = KeyPair::generate().unwrap();
let server_pk = server_keypair.public_key;
let server_sk = server_keypair.secret_key;
let client_keys = client_session_keys(&client_pk, &client_sk, &server_pk).unwrap();
let server_keys = server_session_keys(&server_pk, &server_sk, &client_pk).unwrap();
assert_eq!(client_keys.tx, server_keys.rx);
assert_eq!(client_keys.rx, server_keys.tx);
}
#[test]
fn test_public_key_asref() {
let keypair = KeyPair::generate().unwrap();
let public_key = keypair.public_key;
let bytes_ref: &[u8] = public_key.as_ref();
assert_eq!(bytes_ref.len(), PUBLICKEYBYTES);
assert_eq!(bytes_ref, public_key.as_bytes());
}
#[test]
fn test_public_key_try_from_slice() {
let bytes = [42u8; PUBLICKEYBYTES];
let public_key = PublicKey::try_from(&bytes[..]).unwrap();
assert_eq!(public_key.as_bytes(), &bytes);
let short_bytes = [42u8; PUBLICKEYBYTES - 1];
assert!(PublicKey::try_from(&short_bytes[..]).is_err());
let long_bytes = [42u8; PUBLICKEYBYTES + 1];
assert!(PublicKey::try_from(&long_bytes[..]).is_err());
}
#[test]
fn test_public_key_from_bytes() {
let bytes = [42u8; PUBLICKEYBYTES];
let public_key = PublicKey::from(bytes);
assert_eq!(public_key.as_bytes(), &bytes);
}
#[test]
fn test_public_key_into_bytes() {
let bytes = [42u8; PUBLICKEYBYTES];
let public_key = PublicKey::from(bytes);
let array: [u8; PUBLICKEYBYTES] = public_key.into();
assert_eq!(array, bytes);
}
#[test]
fn test_secret_key_asref() {
let keypair = KeyPair::generate().unwrap();
let secret_key = keypair.secret_key;
let bytes_ref: &[u8] = secret_key.as_ref();
assert_eq!(bytes_ref.len(), SECRETKEYBYTES);
assert_eq!(bytes_ref, secret_key.as_bytes());
}
#[test]
fn test_secret_key_try_from_slice() {
let bytes = [42u8; SECRETKEYBYTES];
let secret_key = SecretKey::try_from(&bytes[..]).unwrap();
assert_eq!(secret_key.as_bytes(), &bytes);
let short_bytes = [42u8; SECRETKEYBYTES - 1];
assert!(SecretKey::try_from(&short_bytes[..]).is_err());
let long_bytes = [42u8; SECRETKEYBYTES + 1];
assert!(SecretKey::try_from(&long_bytes[..]).is_err());
}
#[test]
fn test_secret_key_from_bytes() {
let bytes = [42u8; SECRETKEYBYTES];
let secret_key = SecretKey::from(bytes);
assert_eq!(secret_key.as_bytes(), &bytes);
}
#[test]
fn test_secret_key_into_bytes() {
let bytes = [42u8; SECRETKEYBYTES];
let secret_key = SecretKey::from(bytes);
let array: [u8; SECRETKEYBYTES] = secret_key.into();
assert_eq!(array, bytes);
}
#[test]
fn test_secret_key_zeroization() {
let bytes = [42u8; SECRETKEYBYTES];
let secret_key = SecretKey::from(bytes);
let array: [u8; SECRETKEYBYTES] = secret_key.into();
assert_eq!(array, bytes);
}
#[test]
fn test_roundtrip_conversions() {
let keypair = KeyPair::generate().unwrap();
let pk = keypair.public_key;
let pk_bytes = pk.as_bytes().to_vec();
let pk2 = PublicKey::try_from(&pk_bytes[..]).unwrap();
assert_eq!(pk2.as_bytes(), pk_bytes);
let pk_array: [u8; PUBLICKEYBYTES] = pk2.into();
let pk3 = PublicKey::from(pk_array);
assert_eq!(pk3.as_bytes(), pk_bytes);
let sk = keypair.secret_key;
let sk_bytes = sk.as_bytes().to_vec();
let sk2 = SecretKey::try_from(&sk_bytes[..]).unwrap();
assert_eq!(sk2.as_bytes(), sk_bytes);
let sk_array: [u8; SECRETKEYBYTES] = sk2.into();
let sk3 = SecretKey::from(sk_array);
assert_eq!(sk3.as_bytes(), sk_bytes);
}
}