pub(crate) mod buffers;
pub mod cipher;
pub mod cipher_state;
pub mod curve;
pub mod error;
pub(crate) mod handshake;
pub mod hash;
#[cfg(feature = "async-io")]
#[cfg_attr(docsrs, doc(cfg(feature = "async-io")))]
#[allow(clippy::type_complexity)]
pub mod io_async;
#[allow(clippy::type_complexity)]
pub mod io_sync;
pub mod pattern;
pub(crate) mod process;
pub mod role;
#[cfg(any(target_os = "macos", target_os = "ios", test))]
pub(crate) mod seal;
pub mod session_id;
pub mod symmetric_state;
pub mod tokens;
pub mod transport;
pub mod well_formed;
pub use self::cipher::{ChaChaPoly, Cipher};
pub use self::cipher_state::CipherState;
pub use self::curve::{Curve, DhCurve, P256, X448, X25519};
pub use self::error::HandshakeError;
#[cfg(feature = "async-io")]
#[cfg_attr(docsrs, doc(cfg(feature = "async-io")))]
pub use self::io_async::{AsyncHandshake, AsyncReceiving, AsyncSending, AsyncTransport};
pub use self::io_sync::{SyncHandshake, SyncReceiving, SyncSending, SyncTransport};
pub use self::hash::{Blake2b, Hash};
pub use self::pattern::Pattern;
pub use self::role::{Initiator, Responder, Role};
pub use self::session_id::SessionId;
pub use self::symmetric_state::SymmetricState;
pub use self::tokens::*;
pub use self::transport::{Transport, TransportRecv, TransportSend};
pub use self::well_formed::WellFormed;
use std::fmt;
use std::marker::PhantomData;
pub struct Noise<P, Cu, Ci, H> {
_pattern: PhantomData<fn() -> P>,
_curve: PhantomData<fn() -> Cu>,
_cipher: PhantomData<fn() -> Ci>,
_hash: PhantomData<fn() -> H>,
}
impl<P, Cu, Ci, H> Noise<P, Cu, Ci, H> {
pub const fn new() -> Self {
Self {
_pattern: PhantomData,
_curve: PhantomData,
_cipher: PhantomData,
_hash: PhantomData,
}
}
}
impl<P: Pattern, Cu: DhCurve, Ci: Cipher, H: Hash> Noise<P, Cu, Ci, H> {
pub const DHLEN: usize = Cu::DHLEN;
pub const PUBLIC_KEY_SIZE: usize = Cu::PUBLIC_KEY_SIZE;
pub const TAG_SIZE: usize = Ci::TAG_SIZE;
pub const HASH_LEN: usize = H::HASH_LEN;
pub const NUM_MESSAGES: usize = P::NUM_MESSAGES;
}
impl<P: Pattern, Cu: Curve, Ci: Cipher, H: Hash> fmt::Display for Noise<P, Cu, Ci, H> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "Noise_{}_{}_{}_{}", P::NAME, Cu::NAME, Ci::NAME, H::NAME)
}
}
impl<P, Cu, Ci, H> Default for Noise<P, Cu, Ci, H> {
fn default() -> Self {
Self::new()
}
}
pub trait Protocol {
type Pattern: Pattern;
type Curve: DhCurve;
type Cipher: Cipher;
type Hash: Hash;
}
impl<P: WellFormed, Cu: DhCurve, Ci: Cipher, H: Hash> Protocol for Noise<P, Cu, Ci, H> {
type Pattern = P;
type Curve = Cu;
type Cipher = Ci;
type Hash = H;
}
#[cfg(test)]
mod tests {
use super::*;
use crate::curve::p256::{P256r1PrivateKey, P256r1PublicKey};
use crate::noise_message_size;
use crate::provider::EphemeralOnly;
use crate::provider::ProviderExt;
use crate::psk::Psk;
use rand::{SeedableRng, rngs::StdRng};
use std::cell::RefCell;
use std::collections::VecDeque;
use std::rc::Rc;
type Channel = Noise<pattern::IKpsk1, P256, ChaChaPoly, Blake2b>;
type NoiseSeal = Noise<pattern::N, P256, ChaChaPoly, Blake2b>;
type NoiseK = Noise<pattern::K, P256, ChaChaPoly, Blake2b>;
type NoiseKpsk0 = Noise<pattern::Kpsk0, P256, ChaChaPoly, Blake2b>;
#[derive(Clone)]
struct Pipe {
inbound: Rc<RefCell<VecDeque<u8>>>,
outbound: Rc<RefCell<VecDeque<u8>>>,
}
impl Pipe {
fn pair() -> (Pipe, Pipe) {
let l = Rc::new(RefCell::new(VecDeque::new()));
let r = Rc::new(RefCell::new(VecDeque::new()));
(
Pipe {
inbound: r.clone(),
outbound: l.clone(),
},
Pipe {
inbound: l,
outbound: r,
},
)
}
fn take_written(&self) -> Vec<u8> {
self.outbound.borrow_mut().drain(..).collect()
}
fn feed(&self, bytes: &[u8]) {
self.inbound.borrow_mut().extend(bytes.iter().copied());
}
}
impl std::io::Read for Pipe {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
let mut q = self.inbound.borrow_mut();
let n = q.len().min(buf.len());
for slot in buf.iter_mut().take(n) {
*slot = q.pop_front().unwrap();
}
Ok(n)
}
}
impl std::io::Write for Pipe {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.outbound.borrow_mut().extend(buf.iter().copied());
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
#[test]
fn descriptor_string() {
let proto = Channel::new();
assert_eq!(proto.to_string(), "Noise_IKpsk1_P256_ChaChaPoly_BLAKE2b");
}
#[test]
fn n_descriptor_string() {
let proto = NoiseSeal::new();
assert_eq!(proto.to_string(), "Noise_N_P256_ChaChaPoly_BLAKE2b");
}
#[test]
fn sizes() {
assert_eq!(Channel::DHLEN, 32);
assert_eq!(Channel::PUBLIC_KEY_SIZE, 65);
assert_eq!(Channel::TAG_SIZE, 16);
assert_eq!(Channel::HASH_LEN, 64);
assert_eq!(Channel::NUM_MESSAGES, 2);
}
#[test]
fn n_sizes() {
assert_eq!(NoiseSeal::NUM_MESSAGES, 1);
assert_eq!(size_of::<NoiseSeal>(), 0);
}
#[test]
fn zero_sized() {
assert_eq!(size_of::<Channel>(), 0);
}
#[test]
fn noise_n_seal_open() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let recipient_static = provider.generate::<P256>().unwrap();
let recipient_pub = provider.public(&recipient_static).unwrap();
let psk_to_seal = Psk::from_bytes([0x42; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let sealer = SyncHandshake::<NoiseSeal, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(recipient_pub);
let (mut transport, _) = sealer.e().unwrap().es().unwrap().into_parts();
let msg = i_pipe.take_written();
assert_eq!(msg.len(), 81);
let mut sealed = [0u8; 64]; let sealed_len = transport.send(psk_to_seal.as_bytes(), &mut sealed).unwrap();
r_pipe.feed(&msg);
let opener = SyncHandshake::<NoiseSeal, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(recipient_static)
.unwrap();
let (_, recv) = opener.recv().e().unwrap();
let (mut transport, _) = recv.es().unwrap().into_parts();
let mut opened = [0u8; 32];
let opened_len = transport
.receive(&sealed[..sealed_len], &mut opened)
.unwrap();
assert_eq!(opened_len, 32);
assert_eq!(opened, *psk_to_seal.as_bytes());
}
#[test]
fn noise_n_tampered_ephemeral_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let recipient_static = provider.generate::<P256>().unwrap();
let recipient_pub = provider.public(&recipient_static).unwrap();
let (i_pipe, r_pipe) = Pipe::pair();
let sealer = SyncHandshake::<NoiseSeal, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(recipient_pub);
let (_transport, _) = sealer.e().unwrap().es().unwrap().into_parts();
let mut tampered = i_pipe.take_written();
tampered[1] ^= 0xFF;
r_pipe.feed(&tampered);
let opener = SyncHandshake::<NoiseSeal, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(recipient_static)
.unwrap();
let (_, recv) = match opener.recv().e() {
Err(_) => return, Ok(result) => result,
};
assert!(recv.es().is_err());
}
#[test]
fn noise_n_tampered_tag_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let recipient_static = provider.generate::<P256>().unwrap();
let recipient_pub = provider.public(&recipient_static).unwrap();
let (i_pipe, r_pipe) = Pipe::pair();
let sealer = SyncHandshake::<NoiseSeal, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(recipient_pub);
let (_transport, _) = sealer.e().unwrap().es().unwrap().into_parts();
let mut tampered = i_pipe.take_written();
let len = tampered.len();
tampered[len - 1] ^= 0xFF;
r_pipe.feed(&tampered);
let opener = SyncHandshake::<NoiseSeal, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(recipient_static)
.unwrap();
let (_, recv) = opener.recv().e().unwrap();
assert!(recv.es().is_err());
}
#[test]
fn k_descriptor_string() {
let proto = NoiseK::new();
assert_eq!(proto.to_string(), "Noise_K_P256_ChaChaPoly_BLAKE2b");
}
#[test]
fn kpsk0_descriptor_string() {
let proto = NoiseKpsk0::new();
assert_eq!(proto.to_string(), "Noise_Kpsk0_P256_ChaChaPoly_BLAKE2b");
}
#[test]
fn k_sizes() {
assert_eq!(NoiseK::NUM_MESSAGES, 1);
assert_eq!(size_of::<NoiseK>(), 0);
}
#[test]
fn kpsk0_sizes() {
assert_eq!(NoiseKpsk0::NUM_MESSAGES, 1);
assert_eq!(size_of::<NoiseKpsk0>(), 0);
}
#[test]
fn noise_k_seal_open() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let alice_static = provider.generate::<P256>().unwrap();
let alice_pub = provider.public(&alice_static).unwrap();
let bob_static = provider.generate::<P256>().unwrap();
let bob_pub = provider.public(&bob_static).unwrap();
let payload: [u8; 32] = [0x42; 32];
let (i_pipe, r_pipe) = Pipe::pair();
let sealer = SyncHandshake::<NoiseK, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_s(alice_static)
.unwrap()
.set_rs(bob_pub);
let (mut transport, _) = sealer.e().unwrap().es().unwrap().ss().unwrap().into_parts();
let msg = i_pipe.take_written();
assert_eq!(msg.len(), 81);
let mut sealed = [0u8; 64]; let sealed_len = transport.send(&payload, &mut sealed).unwrap();
r_pipe.feed(&msg);
let opener = SyncHandshake::<NoiseK, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_rs(alice_pub)
.set_s(bob_static)
.unwrap();
let (_, recv) = opener.recv().e().unwrap();
let recv = recv.es().unwrap();
let (mut transport, _) = recv.ss().unwrap().into_parts();
let mut opened = [0u8; 32];
let opened_len = transport
.receive(&sealed[..sealed_len], &mut opened)
.unwrap();
assert_eq!(opened_len, 32);
assert_eq!(opened, payload);
}
#[test]
fn noise_kpsk0_seal_open() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let alice_static = provider.generate::<P256>().unwrap();
let alice_pub = provider.public(&alice_static).unwrap();
let bob_static = provider.generate::<P256>().unwrap();
let bob_pub = provider.public(&bob_static).unwrap();
let psk = Psk::from_bytes([0xBB; 32]);
let payload: [u8; 32] = [0x42; 32];
let (i_pipe, r_pipe) = Pipe::pair();
let sealer = SyncHandshake::<NoiseKpsk0, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_s(alice_static)
.unwrap()
.set_rs(bob_pub);
let (mut transport, _) = sealer
.psk(&psk)
.unwrap()
.e()
.unwrap()
.es()
.unwrap()
.ss()
.unwrap()
.into_parts();
let msg = i_pipe.take_written();
assert_eq!(msg.len(), 81);
let mut sealed = [0u8; 64];
let sealed_len = transport.send(&payload, &mut sealed).unwrap();
r_pipe.feed(&msg);
let opener = SyncHandshake::<NoiseKpsk0, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_rs(alice_pub)
.set_s(bob_static)
.unwrap();
let recv = opener.recv();
let recv = recv.psk(&psk).unwrap();
let (_, recv) = recv.e().unwrap();
let recv = recv.es().unwrap();
let (mut transport, _) = recv.ss().unwrap().into_parts();
let mut opened = [0u8; 32];
let opened_len = transport
.receive(&sealed[..sealed_len], &mut opened)
.unwrap();
assert_eq!(opened_len, 32);
assert_eq!(opened, payload);
}
#[test]
fn noise_kpsk0_wrong_psk_fails() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let alice_static = provider.generate::<P256>().unwrap();
let alice_pub = provider.public(&alice_static).unwrap();
let bob_static = provider.generate::<P256>().unwrap();
let bob_pub = provider.public(&bob_static).unwrap();
let psk = Psk::from_bytes([0xBB; 32]);
let wrong_psk = Psk::from_bytes([0xCC; 32]);
let payload: [u8; 32] = [0x42; 32];
let (i_pipe, r_pipe) = Pipe::pair();
let sealer = SyncHandshake::<NoiseKpsk0, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_s(alice_static)
.unwrap()
.set_rs(bob_pub);
let (mut transport, _) = sealer
.psk(&psk)
.unwrap()
.e()
.unwrap()
.es()
.unwrap()
.ss()
.unwrap()
.into_parts();
let msg = i_pipe.take_written();
let mut sealed = [0u8; 64];
let _sealed_len = transport.send(&payload, &mut sealed).unwrap();
r_pipe.feed(&msg);
let opener = SyncHandshake::<NoiseKpsk0, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_rs(alice_pub)
.set_s(bob_static)
.unwrap();
let recv = opener.recv();
let recv = recv.psk(&wrong_psk).unwrap();
let (_, recv) = recv.e().unwrap();
let recv = recv.es().unwrap();
let result = recv.ss();
assert!(result.is_err());
}
#[test]
fn ikpsk1_round_trip() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let initiator_pub = provider.public(&initiator_static).unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xAA; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let msg1 = i_pipe.take_written();
assert_eq!(msg1.len(), 162);
r_pipe.feed(&msg1);
let initiator_e_from_wire =
P256r1PublicKey::from_bytes(&msg1[..65]).expect("valid ephemeral in msg1");
let (initiator_ephemeral, recv) = r_hs.recv().e().unwrap();
assert_eq!(initiator_ephemeral, initiator_e_from_wire);
let recv = recv.es().unwrap();
let (revealed_initiator_pub, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
assert_eq!(revealed_initiator_pub, initiator_pub);
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let msg2 = r_pipe.take_written();
assert_eq!(msg2.len(), 81);
i_pipe.feed(&msg2);
let responder_e_from_wire =
P256r1PublicKey::from_bytes(&msg2[..65]).expect("valid ephemeral in msg2");
let (responder_ephemeral, recv) = i_hs.recv().e().unwrap();
assert_eq!(responder_ephemeral, responder_e_from_wire);
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
assert_eq!(i_transport.session_id(), r_transport.session_id());
let mut i_transport = i_transport;
let mut r_transport = r_transport;
let plaintext = b"hello from initiator";
let mut ct_buf = [0u8; 256];
let ct_len = i_transport.send(plaintext, &mut ct_buf).unwrap();
let mut pt_buf = [0u8; 256];
let pt_len = r_transport.receive(&ct_buf[..ct_len], &mut pt_buf).unwrap();
assert_eq!(&pt_buf[..pt_len], plaintext);
let plaintext = b"hello from responder";
let ct_len = r_transport.send(plaintext, &mut ct_buf).unwrap();
let pt_len = i_transport.receive(&ct_buf[..ct_len], &mut pt_buf).unwrap();
assert_eq!(&pt_buf[..pt_len], plaintext);
}
#[test]
fn chacha_encrypt_decrypt_round_trip() {
let key = [0x42u8; 32];
let plaintext = b"the quick brown fox";
let ad = b"associated data";
let mut ct = [0u8; 128];
let ct_len = ChaChaPoly::encrypt(&key, 0, ad, plaintext, &mut ct).unwrap();
assert_eq!(ct_len, plaintext.len() + 16);
let mut pt = [0u8; 128];
let pt_len = ChaChaPoly::decrypt(&key, 0, ad, &ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], plaintext);
}
#[test]
fn chacha_decrypt_corrupted_tag() {
let key = [0x42u8; 32];
let plaintext = b"test data";
let mut ct = [0u8; 64];
let ct_len = ChaChaPoly::encrypt(&key, 0, &[], plaintext, &mut ct).unwrap();
ct[ct_len - 1] ^= 0xFF;
let mut pt = [0u8; 64];
let err = ChaChaPoly::decrypt(&key, 0, &[], &ct[..ct_len], &mut pt).unwrap_err();
assert!(
matches!(err, error::HandshakeError::DecryptionFailed),
"expected DecryptionFailed, got {err:?}"
);
}
#[test]
fn chacha_decrypt_too_short() {
let key = [0u8; 32];
let mut pt = [0u8; 64];
let err = ChaChaPoly::decrypt(&key, 0, &[], &[0u8; 15], &mut pt).unwrap_err();
assert!(matches!(err, error::HandshakeError::DecryptionFailed));
}
#[test]
fn chacha_wrong_nonce_fails() {
let key = [0x42u8; 32];
let plaintext = b"nonce matters";
let mut ct = [0u8; 64];
let ct_len = ChaChaPoly::encrypt(&key, 0, &[], plaintext, &mut ct).unwrap();
let mut pt = [0u8; 64];
let err = ChaChaPoly::decrypt(&key, 1, &[], &ct[..ct_len], &mut pt).unwrap_err();
assert!(matches!(err, error::HandshakeError::DecryptionFailed));
}
#[test]
fn chacha_wrong_key_fails() {
let key = [0x42u8; 32];
let wrong_key = [0x43u8; 32];
let plaintext = b"key matters";
let mut ct = [0u8; 64];
let ct_len = ChaChaPoly::encrypt(&key, 0, &[], plaintext, &mut ct).unwrap();
let mut pt = [0u8; 64];
let err = ChaChaPoly::decrypt(&wrong_key, 0, &[], &ct[..ct_len], &mut pt).unwrap_err();
assert!(matches!(err, error::HandshakeError::DecryptionFailed));
}
#[test]
fn chacha_wrong_ad_fails() {
let key = [0x42u8; 32];
let plaintext = b"ad matters";
let mut ct = [0u8; 64];
let ct_len = ChaChaPoly::encrypt(&key, 0, b"correct", plaintext, &mut ct).unwrap();
let mut pt = [0u8; 64];
let err = ChaChaPoly::decrypt(&key, 0, b"wrong", &ct[..ct_len], &mut pt).unwrap_err();
assert!(matches!(err, error::HandshakeError::DecryptionFailed));
}
#[test]
fn chacha_empty_plaintext() {
let key = [0x42u8; 32];
let mut ct = [0u8; 16]; let ct_len = ChaChaPoly::encrypt(&key, 0, &[], &[], &mut ct).unwrap();
assert_eq!(ct_len, 16);
let mut pt = [0u8; 0];
let pt_len = ChaChaPoly::decrypt(&key, 0, &[], &ct[..ct_len], &mut pt).unwrap();
assert_eq!(pt_len, 0);
}
#[test]
fn cipher_state_unkeyed_passthrough() {
let mut cs = cipher_state::CipherState::<ChaChaPoly>::empty();
assert!(!cs.has_key());
let plaintext = b"plaintext passthrough";
let mut out = [0u8; 64];
let len = cs.encrypt_with_ad(b"ad", plaintext, &mut out).unwrap();
assert_eq!(&out[..len], plaintext);
let mut pt = [0u8; 64];
let len = cs.decrypt_with_ad(b"ad", &out[..len], &mut pt).unwrap();
assert_eq!(&pt[..len], plaintext);
}
#[test]
fn blake2b_hash_deterministic() {
let a = Blake2b::hash(b"test data");
let b = Blake2b::hash(b"test data");
assert_eq!(a, b);
assert_eq!(a.len(), 64);
}
#[test]
fn blake2b_hash_different_inputs() {
let a = Blake2b::hash(b"input one");
let b = Blake2b::hash(b"input two");
assert_ne!(a, b);
}
#[test]
fn blake2b_hash_two_equals_concat() {
let a = b"first part";
let b = b"second part";
let h1 = Blake2b::hash_two(a, b);
let mut concat = a.to_vec();
concat.extend_from_slice(b);
let h2 = Blake2b::hash(&concat);
assert_eq!(h1, h2);
}
#[test]
fn blake2b_hmac_deterministic() {
let a = Blake2b::hmac(b"key", b"data");
let b = Blake2b::hmac(b"key", b"data");
assert_eq!(a, b);
assert_eq!(a.len(), 64);
}
#[test]
fn blake2b_hmac_different_keys() {
let a = Blake2b::hmac(b"key1", b"data");
let b = Blake2b::hmac(b"key2", b"data");
assert_ne!(a, b);
}
#[test]
fn wrong_message_length_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let responder_static = provider.generate::<P256>().unwrap();
let (_unused, r_pipe) = Pipe::pair();
let bad_msg = [0u8; 64];
r_pipe.feed(&bad_msg);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
assert!(r_hs.recv().e().is_err());
}
#[test]
fn expected_message_size_reports_correctly() {
assert_eq!(
noise_message_size!(curve: P256, cipher: ChaChaPoly, has_psk: true, keyed: false, tokens: [E, Es, S, Ss, Psk],),
162
);
}
#[test]
fn corrupted_encrypted_static_in_msg1_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xBB; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let _i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let mut corrupted = i_pipe.take_written();
corrupted[70] ^= 0xFF;
r_pipe.feed(&corrupted);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
assert!(recv.s().is_err());
}
#[test]
fn mismatched_psk_fails() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let i_psk = Psk::from_bytes([0xAA; 32]);
let r_psk = Psk::from_bytes([0xBB; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let _i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&i_psk)
.unwrap();
let msg1 = i_pipe.take_written();
r_pipe.feed(&msg1);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let result = recv.psk(&r_psk);
assert!(result.is_err());
}
#[test]
fn transport_corrupted_ciphertext_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xCC; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
let mut i_transport = i_transport;
let mut r_transport = r_transport;
let mut ct_buf = [0u8; 256];
let ct_len = i_transport.send(b"secret", &mut ct_buf).unwrap();
ct_buf[0] ^= 0xFF;
let mut pt_buf = [0u8; 256];
let err = r_transport
.receive(&ct_buf[..ct_len], &mut pt_buf)
.unwrap_err();
assert!(matches!(err, error::HandshakeError::DecryptionFailed));
}
#[test]
fn transport_multiple_messages_nonce_advances() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xDD; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
let mut i_transport = i_transport;
let mut r_transport = r_transport;
let mut ct_buf = [0u8; 256];
let mut pt_buf = [0u8; 256];
for i in 0u32..10 {
let msg = format!("message {i}");
let ct_len = i_transport.send(msg.as_bytes(), &mut ct_buf).unwrap();
let pt_len = r_transport.receive(&ct_buf[..ct_len], &mut pt_buf).unwrap();
assert_eq!(&pt_buf[..pt_len], msg.as_bytes());
let reply = format!("reply {i}");
let ct_len = r_transport.send(reply.as_bytes(), &mut ct_buf).unwrap();
let pt_len = i_transport.receive(&ct_buf[..ct_len], &mut pt_buf).unwrap();
assert_eq!(&pt_buf[..pt_len], reply.as_bytes());
}
}
#[test]
fn symmetric_state_long_protocol_name() {
let long_name = "A".repeat(100);
let ss = symmetric_state::SymmetricState::<ChaChaPoly, Blake2b>::initialize(&long_name);
assert_eq!(ss.handshake_hash().len(), 64);
}
#[test]
fn ikpsk1_wrong_responder_key_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let wrong_static = provider.generate::<P256>().unwrap();
let wrong_pub = provider.public(&wrong_static).unwrap();
let psk = Psk::from_bytes([0xCC; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(wrong_pub);
let _i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let msg1 = i_pipe.take_written();
r_pipe.feed(&msg1);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let result = recv.s();
assert!(result.is_err());
}
#[test]
fn transport_keys_are_directional() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xFF; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
let mut i_transport = i_transport;
let mut ct_buf = [0u8; 256];
let ct_len = i_transport.send(b"hello", &mut ct_buf).unwrap();
let mut pt_buf = [0u8; 256];
let err = i_transport
.receive(&ct_buf[..ct_len], &mut pt_buf)
.unwrap_err();
assert!(matches!(err, error::HandshakeError::DecryptionFailed));
let mut r_transport = r_transport;
let pt_len = r_transport.receive(&ct_buf[..ct_len], &mut pt_buf).unwrap();
assert_eq!(&pt_buf[..pt_len], b"hello");
}
#[test]
fn two_sessions_produce_different_handshake_hashes() {
let responder_bytes = [0xBB_u8; 32];
let psk = Psk::from_bytes([0xAA; 32]);
let mut hashes = Vec::new();
for _ in 0..2 {
let responder_static =
P256r1PrivateKey::from_bytes(responder_bytes).expect("valid test scalar");
let responder_pub = responder_static.public();
let initiator_static = EphemeralOnly::new(StdRng::from_os_rng())
.generate::<P256>()
.unwrap();
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
assert_eq!(i_transport.session_id(), r_transport.session_id());
hashes.push(i_transport.session_id().as_ref().to_vec());
}
assert_ne!(hashes[0], hashes[1]);
}
#[tokio::test]
async fn ikpsk1_wrong_initiator_static_in_msg1() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let initiator_pub = provider.public(&initiator_static).unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xCC; 32]);
let wrong_static = provider.generate::<P256>().unwrap();
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(wrong_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (revealed_pub, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
assert_ne!(revealed_pub, initiator_pub);
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
assert_eq!(i_transport.session_id(), r_transport.session_id());
drop(i_transport);
drop(r_transport);
}
#[test]
fn ikpsk1_corrupted_msg1_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xDD; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let _i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let mut corrupted = i_pipe.take_written();
corrupted[5] ^= 0xFF;
r_pipe.feed(&corrupted);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
match r_hs.recv().e() {
Err(_) => {} Ok((_, recv)) => {
assert!(recv.es().is_err());
}
}
}
#[test]
fn ikpsk1_corrupted_msg2_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xDD; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (_r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let mut corrupted = r_pipe.take_written();
corrupted[3] ^= 0xFF;
i_pipe.feed(&corrupted);
match i_hs.recv().e() {
Err(_) => {} Ok((_, recv)) => {
assert!(recv.ee().is_err());
}
}
}
#[test]
fn transport_replayed_message_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xEE; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
let mut i_transport = i_transport;
let mut r_transport = r_transport;
let mut ct_buf = [0u8; 256];
let ct_len = i_transport.send(b"first message", &mut ct_buf).unwrap();
let captured = ct_buf[..ct_len].to_vec();
let mut pt_buf = [0u8; 256];
let pt_len = r_transport.receive(&captured, &mut pt_buf).unwrap();
assert_eq!(&pt_buf[..pt_len], b"first message");
let err = r_transport.receive(&captured, &mut pt_buf).unwrap_err();
assert!(matches!(err, error::HandshakeError::DecryptionFailed));
}
#[test]
fn transport_enforces_max_message_length() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0xFF; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
let mut i_transport = i_transport;
let mut r_transport = r_transport;
let max_payload: Vec<u8> = (0..65519).map(|i| (i % 256) as u8).collect();
let mut ct_buf = vec![0u8; max_payload.len() + 16];
let ct_len = i_transport.send(&max_payload, &mut ct_buf).unwrap();
assert_eq!(ct_len, 65535);
let mut pt_buf = vec![0u8; max_payload.len()];
let pt_len = r_transport.receive(&ct_buf[..ct_len], &mut pt_buf).unwrap();
assert_eq!(&pt_buf[..pt_len], &max_payload[..]);
let over_payload = vec![0u8; 65520];
let mut ct_buf = vec![0u8; over_payload.len() + 16];
let err = i_transport.send(&over_payload, &mut ct_buf).unwrap_err();
assert!(matches!(
err,
error::HandshakeError::MessageTooLong { len: 65536 }
));
let oversize = vec![0u8; 65536];
let mut pt_buf = vec![0u8; oversize.len()];
let err = r_transport.receive(&oversize, &mut pt_buf).unwrap_err();
assert!(matches!(
err,
error::HandshakeError::MessageTooLong { len: 65536 }
));
}
#[test]
fn transport_rekey_then_communicate() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0x11; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
let mut i_transport = i_transport;
let mut r_transport = r_transport;
let mut ct = [0u8; 256];
let mut pt = [0u8; 256];
let ct_len = i_transport.send(b"before rekey", &mut ct).unwrap();
let pt_len = r_transport.receive(&ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], b"before rekey");
i_transport.rekey().unwrap();
r_transport.rekey().unwrap();
let ct_len = i_transport.send(b"after rekey", &mut ct).unwrap();
let pt_len = r_transport.receive(&ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], b"after rekey");
let ct_len = r_transport.send(b"reply after rekey", &mut ct).unwrap();
let pt_len = i_transport.receive(&ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], b"reply after rekey");
}
#[test]
fn transport_split_round_trip() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0x5A; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
assert!(i_transport.local_ephemeral().is_some());
assert!(i_transport.remote_ephemeral().is_some());
assert!(r_transport.local_ephemeral().is_some());
assert!(r_transport.remote_ephemeral().is_some());
assert!(i_transport.session_id() == r_transport.session_id());
let (mut i_send, mut i_recv) = i_transport.split();
let (mut r_send, mut r_recv) = r_transport.split();
assert!(i_send.session_id() == i_recv.session_id());
assert!(i_send.session_id() == r_recv.session_id());
assert!(i_send.local_ephemeral().is_some());
assert!(i_send.remote_ephemeral().is_some());
assert!(r_recv.local_ephemeral().is_some());
assert!(r_recv.remote_ephemeral().is_some());
let mut ct = [0u8; 256];
let mut pt = [0u8; 256];
let ct_len = i_send.encrypt(b"ping", &mut ct).unwrap();
assert_eq!(ct_len, b"ping".len() + TransportSend::<Channel>::OVERHEAD);
let pt_len = r_recv.decrypt(&ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], b"ping");
let ct_len = r_send.encrypt(b"pong", &mut ct).unwrap();
let pt_len = i_recv.decrypt(&ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], b"pong");
i_send.rekey().unwrap();
r_recv.rekey().unwrap();
r_send.rekey().unwrap();
i_recv.rekey().unwrap();
let ct_len = i_send.encrypt(b"after rekey", &mut ct).unwrap();
let pt_len = r_recv.decrypt(&ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], b"after rekey");
let ct_len = r_send.encrypt(b"reply after rekey", &mut ct).unwrap();
let pt_len = i_recv.decrypt(&ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], b"reply after rekey");
}
#[test]
fn transport_rekey_desync_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let initiator_static = provider.generate::<P256>().unwrap();
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let psk = Psk::from_bytes([0x22; 32]);
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
let mut i_transport = i_transport;
let mut r_transport = r_transport;
i_transport.rekey().unwrap();
let mut ct = [0u8; 256];
let ct_len = i_transport.send(b"desynced", &mut ct).unwrap();
let mut pt = [0u8; 256];
let err = r_transport.receive(&ct[..ct_len], &mut pt).unwrap_err();
assert!(matches!(err, error::HandshakeError::DecryptionFailed));
}
#[test]
fn rekey_without_key_fails() {
let mut cs = cipher_state::CipherState::<ChaChaPoly>::empty();
let err = cs.rekey().unwrap_err();
assert!(matches!(err, error::HandshakeError::RekeyWithoutKey));
}
#[test]
fn matching_prologue_succeeds() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
let prologue = b"hiss/v1";
type NoiseSeal = Noise<pattern::N, P256, ChaChaPoly, Blake2b>;
let (i_pipe, r_pipe) = Pipe::pair();
let sealer = SyncHandshake::<NoiseSeal, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
prologue,
i_pipe.clone(),
)
.set_rs(responder_pub);
let (mut i_transport, _) = sealer.e().unwrap().es().unwrap().into_parts();
let msg = i_pipe.take_written();
r_pipe.feed(&msg);
let opener = SyncHandshake::<NoiseSeal, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
prologue,
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let (_, recv) = opener.recv().e().unwrap();
let (mut r_transport, _) = recv.es().unwrap().into_parts();
let mut ct = [0u8; 64];
let mut pt = [0u8; 64];
let ct_len = i_transport.send(b"hello", &mut ct).unwrap();
let pt_len = r_transport.receive(&ct[..ct_len], &mut pt).unwrap();
assert_eq!(&pt[..pt_len], b"hello");
}
#[test]
fn mismatched_prologue_rejected() {
let mut provider = EphemeralOnly::new(StdRng::from_os_rng());
let responder_static = provider.generate::<P256>().unwrap();
let responder_pub = provider.public(&responder_static).unwrap();
type NoiseSeal = Noise<pattern::N, P256, ChaChaPoly, Blake2b>;
let (i_pipe, r_pipe) = Pipe::pair();
let sealer = SyncHandshake::<NoiseSeal, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
b"v1",
i_pipe.clone(),
)
.set_rs(responder_pub);
let (_i_transport, _) = sealer.e().unwrap().es().unwrap().into_parts();
let msg = i_pipe.take_written();
r_pipe.feed(&msg);
let opener = SyncHandshake::<NoiseSeal, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
b"v2",
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let (_, recv) = opener.recv().e().unwrap();
assert!(recv.es().is_err());
}
#[test]
fn chacha_encrypt_output_buffer_too_small() {
let key = [0u8; 32];
let plaintext = b"hello world";
let mut output = [0u8; 10]; let err = ChaChaPoly::encrypt(&key, 0, &[], plaintext, &mut output).unwrap_err();
assert!(matches!(
err,
error::HandshakeError::OutputBufferTooSmall { .. }
));
}
#[test]
fn chacha_decrypt_output_buffer_too_small() {
let key = [0u8; 32];
let plaintext = b"hello world";
let mut ct = [0u8; 64];
let ct_len = ChaChaPoly::encrypt(&key, 0, &[], plaintext, &mut ct).unwrap();
let mut output = [0u8; 5]; let err = ChaChaPoly::decrypt(&key, 0, &[], &ct[..ct_len], &mut output).unwrap_err();
assert!(matches!(
err,
error::HandshakeError::OutputBufferTooSmall { .. }
));
}
mod prop {
use super::*;
use crate::psk::Psk;
use proptest::prelude::*;
fn full_ikpsk1_handshake(
initiator_static: P256r1PrivateKey,
responder_static: P256r1PrivateKey,
psk: Psk,
) -> (transport::Transport<Channel>, transport::Transport<Channel>) {
let provider = EphemeralOnly::new(StdRng::from_os_rng());
let responder_pub = provider.public(&responder_static).unwrap();
let (i_pipe, r_pipe) = Pipe::pair();
let i_hs = SyncHandshake::<Channel, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
i_pipe.clone(),
)
.set_rs(responder_pub);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(responder_static)
.unwrap();
let i_hs = i_hs
.e()
.unwrap()
.es()
.unwrap()
.s(initiator_static)
.unwrap()
.ss()
.unwrap()
.psk(&psk)
.unwrap();
let (_, recv) = r_hs.recv().e().unwrap();
let recv = recv.es().unwrap();
let (_, recv) = recv.s().unwrap();
let recv = recv.ss().unwrap();
let r_hs = recv.psk(&psk).unwrap();
let (r_transport, _) = r_hs.e().unwrap().ee().unwrap().se().unwrap().into_parts();
let (_, recv) = i_hs.recv().e().unwrap();
let (i_transport, _) = recv.ee().unwrap().se().unwrap().into_parts();
(i_transport, r_transport)
}
proptest! {
#[test]
fn transport_any_payload(
plaintext in proptest::collection::vec(any::<u8>(), 0..4096),
psk in any::<[u8; 32]>().prop_map(Psk::from_bytes),
) {
let i_sk = EphemeralOnly::new(StdRng::from_os_rng()).generate::<P256>().unwrap();
let r_sk = EphemeralOnly::new(StdRng::from_os_rng()).generate::<P256>().unwrap();
let (mut i_t, mut r_t) = full_ikpsk1_handshake(i_sk, r_sk, psk);
let mut ct = vec![0u8; plaintext.len() + 16];
let ct_len = i_t.send(&plaintext, &mut ct).unwrap();
let mut pt = vec![0u8; plaintext.len()];
let pt_len = r_t.receive(&ct[..ct_len], &mut pt).unwrap();
prop_assert_eq!(&pt[..pt_len], &plaintext[..]);
let ct_len = r_t.send(&plaintext, &mut ct).unwrap();
let pt_len = i_t.receive(&ct[..ct_len], &mut pt).unwrap();
prop_assert_eq!(&pt[..pt_len], &plaintext[..]);
}
#[test]
fn transport_any_corruption_detected(
plaintext in proptest::collection::vec(any::<u8>(), 1..512),
psk in any::<[u8; 32]>().prop_map(Psk::from_bytes),
corrupt_pos_seed in any::<usize>(),
) {
let i_sk = EphemeralOnly::new(StdRng::from_os_rng()).generate::<P256>().unwrap();
let r_sk = EphemeralOnly::new(StdRng::from_os_rng()).generate::<P256>().unwrap();
let (mut i_t, mut r_t) = full_ikpsk1_handshake(i_sk, r_sk, psk);
let mut ct = vec![0u8; plaintext.len() + 16];
let ct_len = i_t.send(&plaintext, &mut ct).unwrap();
let pos = corrupt_pos_seed % ct_len;
ct[pos] ^= 0x01;
let mut pt = vec![0u8; plaintext.len()];
let result = r_t.receive(&ct[..ct_len], &mut pt);
prop_assert!(result.is_err());
}
#[test]
fn random_msg1_rejected(
garbage in proptest::collection::vec(any::<u8>(), 162..163),
) {
let r_sk = EphemeralOnly::new(StdRng::from_os_rng()).generate::<P256>().unwrap();
let (_unused, r_pipe) = Pipe::pair();
r_pipe.feed(&garbage);
let r_hs = SyncHandshake::<Channel, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
r_pipe.clone(),
)
.set_s(r_sk)
.unwrap();
match r_hs.recv().e() {
Err(_) => {} Ok((_, recv)) => {
let recv = recv.es().unwrap();
prop_assert!(recv.s().is_err());
}
}
}
#[test]
fn chacha_round_trip_any(
key in any::<[u8; 32]>(),
nonce in 0..1000u64,
ad in proptest::collection::vec(any::<u8>(), 0..128),
plaintext in proptest::collection::vec(any::<u8>(), 0..4096),
) {
let mut ct = vec![0u8; plaintext.len() + 16];
let ct_len = ChaChaPoly::encrypt(&key, nonce, &ad, &plaintext, &mut ct).unwrap();
let mut pt = vec![0u8; plaintext.len()];
let pt_len = ChaChaPoly::decrypt(&key, nonce, &ad, &ct[..ct_len], &mut pt).unwrap();
prop_assert_eq!(&pt[..pt_len], &plaintext[..]);
}
#[test]
fn chacha_any_bit_flip_detected(
key in any::<[u8; 32]>(),
plaintext in proptest::collection::vec(any::<u8>(), 1..256),
flip_pos_seed in any::<usize>(),
flip_bit in 0u8..8,
) {
let mut ct = vec![0u8; plaintext.len() + 16];
let ct_len = ChaChaPoly::encrypt(&key, 0, &[], &plaintext, &mut ct).unwrap();
let pos = flip_pos_seed % ct_len;
ct[pos] ^= 1 << flip_bit;
let mut pt = vec![0u8; plaintext.len()];
let result = ChaChaPoly::decrypt(&key, 0, &[], &ct[..ct_len], &mut pt);
prop_assert!(result.is_err());
}
}
}
mod single_e_send_finalizer_tests {
use super::super::SyncHandshake;
use super::super::tokens::{Cons, E, Ee, Message, Nil, ToInitiator, ToResponder};
use super::super::{Blake2b, ChaChaPoly, Noise, P256};
use super::super::{Initiator, Pattern, Responder};
use crate::provider::EphemeralOnly;
use rand::{SeedableRng, rngs::StdRng};
struct EThenEe;
impl Pattern for EThenEe {
const NAME: &'static str = "EThenEe";
const NUM_MESSAGES: usize = 2;
type PreMessages = Nil;
type Messages = Cons<
Message<ToResponder, Cons<E, Nil>>,
Cons<Message<ToInitiator, Cons<E, Cons<Ee, Nil>>>, Nil>,
>;
}
type EThenEeProto = Noise<EThenEe, P256, ChaChaPoly, Blake2b>;
#[test]
fn single_e_more_messages_advances_handshake() {
let (i2r, r2i) = (
std::rc::Rc::new(std::cell::RefCell::new(
std::collections::VecDeque::<u8>::new(),
)),
std::rc::Rc::new(std::cell::RefCell::new(
std::collections::VecDeque::<u8>::new(),
)),
);
let init_stream = Pipe {
inbound: r2i.clone(),
outbound: i2r.clone(),
};
let resp_stream = Pipe {
inbound: i2r.clone(),
outbound: r2i.clone(),
};
let initiator = SyncHandshake::<EThenEeProto, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
init_stream,
);
let initiator = initiator.e().unwrap();
assert_eq!(
i2r.borrow().len(),
65,
"a bare `-> e` must be exactly 65 bytes"
);
let responder = SyncHandshake::<EThenEeProto, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
resp_stream,
);
let (_revealed_e, responder) = responder.recv().e().unwrap();
let mut responder_transport = responder.e().unwrap().ee().unwrap();
assert_eq!(
r2i.borrow().len(),
81,
"`<- e, ee` must be exactly 81 bytes"
);
let (_revealed_e, initiator) = initiator.recv().e().unwrap();
let mut initiator_transport = initiator.ee().unwrap();
assert_eq!(
initiator_transport.transport().session_id(),
responder_transport.transport().session_id(),
"initiator and responder must derive a matching session",
);
}
#[derive(Clone)]
struct Pipe {
inbound: std::rc::Rc<std::cell::RefCell<std::collections::VecDeque<u8>>>,
outbound: std::rc::Rc<std::cell::RefCell<std::collections::VecDeque<u8>>>,
}
impl std::io::Read for Pipe {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
let mut q = self.inbound.borrow_mut();
let n = q.len().min(buf.len());
for slot in buf.iter_mut().take(n) {
*slot = q.pop_front().unwrap();
}
Ok(n)
}
}
impl std::io::Write for Pipe {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.outbound.borrow_mut().extend(buf.iter().copied());
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
#[test]
fn single_e_more_messages_sync_streaming() {
let (i2r, r2i) = (
std::rc::Rc::new(std::cell::RefCell::new(
std::collections::VecDeque::<u8>::new(),
)),
std::rc::Rc::new(std::cell::RefCell::new(
std::collections::VecDeque::<u8>::new(),
)),
);
let init_stream = Pipe {
inbound: r2i.clone(),
outbound: i2r.clone(),
};
let resp_stream = Pipe {
inbound: i2r.clone(),
outbound: r2i.clone(),
};
let initiator = SyncHandshake::<EThenEeProto, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
init_stream,
);
let initiator = initiator.e().unwrap();
assert_eq!(
i2r.borrow().len(),
65,
"a bare `-> e` must be exactly 65 bytes"
);
let responder = SyncHandshake::<EThenEeProto, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
resp_stream,
);
let (_revealed_e, responder) = responder.recv().e().unwrap();
let mut responder_transport = responder.e().unwrap().ee().unwrap();
assert_eq!(
r2i.borrow().len(),
81,
"`<- e, ee` must be exactly 81 bytes"
);
let (_revealed_e, initiator) = initiator.recv().e().unwrap();
let mut initiator_transport = initiator.ee().unwrap();
assert_eq!(
initiator_transport.transport().session_id(),
responder_transport.transport().session_id(),
"initiator and responder must derive a matching session",
);
}
#[cfg(feature = "async-io")]
#[tokio::test]
async fn single_e_more_messages_async_streaming() {
use super::super::AsyncHandshake;
let (init_stream, resp_stream) = tokio::io::duplex(4096);
let initiator = AsyncHandshake::<EThenEeProto, Initiator, _, _, _, _>::initiate(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
init_stream,
);
let initiator = initiator.e().await.unwrap();
let responder = AsyncHandshake::<EThenEeProto, Responder, _, _, _, _>::respond(
EphemeralOnly::new(StdRng::from_os_rng()),
&[],
resp_stream,
);
let (_revealed_e, responder) = responder.recv().e().await.unwrap();
let mut responder_transport = responder.e().await.unwrap().ee().await.unwrap();
let (_revealed_e, initiator) = initiator.recv().e().await.unwrap();
let mut initiator_transport = initiator.ee().await.unwrap();
assert_eq!(
initiator_transport.transport().session_id(),
responder_transport.transport().session_id(),
"initiator and responder must derive a matching session",
);
}
}
}