use crate::core::{CryptoError, HASH_SIZE, PUBLIC_KEY_SIZE};
use snow::{Builder, HandshakeState};
use zeroize::Zeroize;
use super::{SessionKey, StaticKeypair, SESSION_KEY_SIZE};
const NOISE_PATTERN: &str = "Noise_IK_25519_ChaChaPoly_BLAKE2s";
pub struct HandshakeResult {
pub handshake_hash: [u8; HASH_SIZE],
}
pub struct InitiatorHandshake {
state: HandshakeState,
}
impl InitiatorHandshake {
pub fn new(
local_keypair: &StaticKeypair,
remote_public: &[u8; PUBLIC_KEY_SIZE],
) -> Result<Self, CryptoError> {
let builder = Builder::new(NOISE_PATTERN.parse().unwrap());
let state = builder
.local_private_key(local_keypair.private_key())
.remote_public_key(remote_public)
.build_initiator()
.map_err(|e| CryptoError::HandshakeFailed(e.to_string()))?;
Ok(Self { state })
}
pub fn write_message(&mut self, payload: &[u8]) -> Result<Vec<u8>, CryptoError> {
let mut buf = vec![0u8; 65535];
let len = self
.state
.write_message(payload, &mut buf)
.map_err(|e| CryptoError::HandshakeFailed(e.to_string()))?;
buf.truncate(len);
Ok(buf)
}
pub fn read_message(mut self, message: &[u8]) -> Result<(Vec<u8>, HandshakeResult), CryptoError> {
let mut payload = vec![0u8; 65535];
let len = self
.state
.read_message(message, &mut payload)
.map_err(|e| CryptoError::HandshakeFailed(e.to_string()))?;
payload.truncate(len);
let hash_slice = self.state.get_handshake_hash();
let mut handshake_hash = [0u8; HASH_SIZE];
handshake_hash.copy_from_slice(hash_slice);
let _transport = self
.state
.into_transport_mode()
.map_err(|e| CryptoError::HandshakeFailed(e.to_string()))?;
Ok((payload, HandshakeResult { handshake_hash }))
}
}
pub struct ResponderHandshake {
state: HandshakeState,
}
impl ResponderHandshake {
pub fn new(local_keypair: &StaticKeypair) -> Result<Self, CryptoError> {
let builder = Builder::new(NOISE_PATTERN.parse().unwrap());
let state = builder
.local_private_key(local_keypair.private_key())
.build_responder()
.map_err(|e| CryptoError::HandshakeFailed(e.to_string()))?;
Ok(Self { state })
}
pub fn read_message(&mut self, message: &[u8]) -> Result<(Vec<u8>, [u8; PUBLIC_KEY_SIZE]), CryptoError> {
let mut payload = vec![0u8; 65535];
let len = self
.state
.read_message(message, &mut payload)
.map_err(|e| CryptoError::HandshakeFailed(e.to_string()))?;
payload.truncate(len);
let remote_static = self
.state
.get_remote_static()
.ok_or_else(|| CryptoError::HandshakeFailed("no remote static key".into()))?;
let mut remote_public = [0u8; PUBLIC_KEY_SIZE];
remote_public.copy_from_slice(remote_static);
Ok((payload, remote_public))
}
pub fn write_message(mut self, payload: &[u8]) -> Result<(Vec<u8>, HandshakeResult), CryptoError> {
let mut buf = vec![0u8; 65535];
let len = self
.state
.write_message(payload, &mut buf)
.map_err(|e| CryptoError::HandshakeFailed(e.to_string()))?;
buf.truncate(len);
let hash_slice = self.state.get_handshake_hash();
let mut handshake_hash = [0u8; HASH_SIZE];
handshake_hash.copy_from_slice(hash_slice);
let _transport = self
.state
.into_transport_mode()
.map_err(|e| CryptoError::HandshakeFailed(e.to_string()))?;
Ok((buf, HandshakeResult { handshake_hash }))
}
}
pub struct SessionKeys {
pub initiator_key: SessionKey,
pub responder_key: SessionKey,
pub handshake_hash: [u8; HASH_SIZE],
pub rekey_auth_key: [u8; HASH_SIZE],
}
impl SessionKeys {
pub fn derive(result: &HandshakeResult, static_dh_secret: &[u8; 32]) -> Result<Self, CryptoError> {
use hkdf::Hkdf;
use sha2::Sha256;
use super::rekey::derive_rekey_auth_key;
let handshake_hash = &result.handshake_hash;
let label = b"nomad v1 session keys";
let hk = Hkdf::<Sha256>::from_prk(handshake_hash)
.map_err(|_| CryptoError::KeyDerivationFailed)?;
let mut key_material = [0u8; 64];
hk.expand(label, &mut key_material)
.map_err(|_| CryptoError::KeyDerivationFailed)?;
let mut initiator_key = [0u8; SESSION_KEY_SIZE];
let mut responder_key = [0u8; SESSION_KEY_SIZE];
initiator_key.copy_from_slice(&key_material[..32]);
responder_key.copy_from_slice(&key_material[32..]);
let rekey_auth_key = derive_rekey_auth_key(static_dh_secret);
key_material.zeroize();
Ok(Self {
initiator_key: SessionKey::from_bytes(initiator_key),
responder_key: SessionKey::from_bytes(responder_key),
handshake_hash: *handshake_hash,
rekey_auth_key,
})
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Role {
Initiator,
Responder,
}
impl SessionKeys {
pub fn send_key(&self, role: Role) -> &SessionKey {
match role {
Role::Initiator => &self.initiator_key,
Role::Responder => &self.responder_key,
}
}
pub fn recv_key(&self, role: Role) -> &SessionKey {
match role {
Role::Initiator => &self.responder_key,
Role::Responder => &self.initiator_key,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_handshake_roundtrip() {
let initiator_keypair = StaticKeypair::generate();
let responder_keypair = StaticKeypair::generate();
let mut initiator = InitiatorHandshake::new(
&initiator_keypair,
responder_keypair.public_key(),
).unwrap();
let mut responder = ResponderHandshake::new(&responder_keypair).unwrap();
let init_payload = b"nomad.echo.v1";
let init_message = initiator.write_message(init_payload).unwrap();
let (recv_payload, remote_public) = responder.read_message(&init_message).unwrap();
assert_eq!(recv_payload, init_payload);
assert_eq!(&remote_public, initiator_keypair.public_key());
let resp_payload = b"OK";
let (resp_message, responder_result) = responder.write_message(resp_payload).unwrap();
let (recv_resp_payload, initiator_result) = initiator.read_message(&resp_message).unwrap();
assert_eq!(recv_resp_payload, resp_payload);
assert_eq!(initiator_result.handshake_hash, responder_result.handshake_hash);
let initiator_static_dh = initiator_keypair.compute_static_dh(responder_keypair.public_key());
let responder_static_dh = responder_keypair.compute_static_dh(initiator_keypair.public_key());
assert_eq!(initiator_static_dh, responder_static_dh);
let initiator_keys = SessionKeys::derive(&initiator_result, &initiator_static_dh).unwrap();
let responder_keys = SessionKeys::derive(&responder_result, &responder_static_dh).unwrap();
assert_eq!(
initiator_keys.send_key(Role::Initiator).as_bytes(),
responder_keys.recv_key(Role::Responder).as_bytes()
);
assert_eq!(
initiator_keys.recv_key(Role::Initiator).as_bytes(),
responder_keys.send_key(Role::Responder).as_bytes()
);
assert_eq!(initiator_keys.rekey_auth_key, responder_keys.rekey_auth_key);
}
#[test]
fn test_handshake_wrong_key_fails() {
let initiator_keypair = StaticKeypair::generate();
let responder_keypair = StaticKeypair::generate();
let wrong_keypair = StaticKeypair::generate();
let mut initiator = InitiatorHandshake::new(
&initiator_keypair,
wrong_keypair.public_key(), ).unwrap();
let mut responder = ResponderHandshake::new(&responder_keypair).unwrap();
let init_message = initiator.write_message(b"test").unwrap();
let result = responder.read_message(&init_message);
assert!(result.is_err());
}
#[test]
fn test_role_keys() {
let initiator_keypair = StaticKeypair::generate();
let responder_keypair = StaticKeypair::generate();
let mut initiator = InitiatorHandshake::new(
&initiator_keypair,
responder_keypair.public_key(),
).unwrap();
let mut responder = ResponderHandshake::new(&responder_keypair).unwrap();
let init_message = initiator.write_message(b"").unwrap();
responder.read_message(&init_message).unwrap();
let (resp_message, responder_result) = responder.write_message(b"").unwrap();
let (_, initiator_result) = initiator.read_message(&resp_message).unwrap();
let static_dh = initiator_keypair.compute_static_dh(responder_keypair.public_key());
let initiator_keys = SessionKeys::derive(&initiator_result, &static_dh).unwrap();
let responder_keys = SessionKeys::derive(&responder_result, &static_dh).unwrap();
assert_eq!(
initiator_keys.send_key(Role::Initiator).as_bytes(),
initiator_keys.initiator_key.as_bytes()
);
assert_eq!(
initiator_keys.recv_key(Role::Initiator).as_bytes(),
initiator_keys.responder_key.as_bytes()
);
assert_eq!(
responder_keys.send_key(Role::Responder).as_bytes(),
responder_keys.responder_key.as_bytes()
);
assert_eq!(
responder_keys.recv_key(Role::Responder).as_bytes(),
responder_keys.initiator_key.as_bytes()
);
}
}