use bytes::BytesMut;
use snow::{Builder, HandshakeState, TransportState, params::NoiseParams};
use std::io;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio_util::codec::Framed;
pub const NOISE_PARAMS: &str = "Noise_XX_25519_ChaChaPoly_BLAKE2s";
pub const NOISE_MAGIC: [u8; 4] = [0xAE, 0x1E, 0x00, 0xFF];
pub const MAX_NOISE_MSG_LEN: usize = 65535;
const LENGTH_PREFIX_LEN: usize = 4;
const NOISE_OVERHEAD: usize = DH_LEN + (DH_LEN + TAG_LEN) + TAG_LEN;
const TAG_LEN: usize = 16;
pub const HANDSHAKE_HASH_LEN: usize = 32;
const MAX_HANDSHAKE_MSG_LEN: usize = 1024;
const _: () = assert!(MAX_HANDSHAKE_MSG_LEN < MAX_NOISE_MSG_LEN);
const DH_LEN: usize = 32;
const FIRST_MESSAGE_PAYLOAD_OFFSET: usize = DH_LEN;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Role {
Initiator,
Responder,
}
impl Role {
const fn tag(&self) -> &'static [u8] {
match self {
Self::Initiator => b"|initiator",
Self::Responder => b"|responder",
}
}
}
pub fn binding_message(domain: &[u8], role: Role, handshake_hash: &[u8]) -> Vec<u8> {
let mut message = Vec::with_capacity(domain.len() + role.tag().len() + handshake_hash.len());
message.extend_from_slice(domain);
message.extend_from_slice(role.tag());
message.extend_from_slice(handshake_hash);
message
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum HandshakeProtocol {
Noise,
Legacy,
}
pub async fn write_noise_magic<S: AsyncWrite + Unpin>(stream: &mut S) -> io::Result<()> {
stream.write_all(&NOISE_MAGIC).await
}
pub async fn detect_handshake_protocol<S: AsyncRead + Unpin>(
stream: &mut S,
) -> io::Result<(HandshakeProtocol, BytesMut)> {
let mut prefix = [0u8; NOISE_MAGIC.len()];
stream.read_exact(&mut prefix).await?;
if prefix == NOISE_MAGIC {
Ok((HandshakeProtocol::Noise, BytesMut::new()))
} else {
Ok((HandshakeProtocol::Legacy, BytesMut::from(&prefix[..])))
}
}
pub fn prepare_framed<S: AsyncRead + AsyncWrite, C>(stream: S, codec: C, read_buf: &[u8]) -> Framed<S, C> {
let mut framed = Framed::new(stream, codec);
framed.read_buffer_mut().extend_from_slice(read_buf);
framed
}
async fn read_frame<S: AsyncRead + Unpin>(stream: &mut S, max_len: usize) -> io::Result<Vec<u8>> {
let mut length = [0u8; LENGTH_PREFIX_LEN];
stream.read_exact(&mut length).await?;
let length = u32::from_le_bytes(length) as usize;
if length > max_len {
return Err(invalid_data(format!("the Noise message is too large ({length} bytes)")));
}
let mut message = vec![0u8; length];
stream.read_exact(&mut message).await?;
Ok(message)
}
async fn write_frame<S: AsyncWrite + Unpin>(stream: &mut S, message: &[u8], max_len: usize) -> io::Result<()> {
if message.len() > max_len {
return Err(invalid_data(format!("the Noise message is too large ({} bytes)", message.len())));
}
let mut framed = Vec::with_capacity(LENGTH_PREFIX_LEN + message.len());
framed.extend_from_slice(&(message.len() as u32).to_le_bytes());
framed.extend_from_slice(message);
stream.write_all(&framed).await?;
stream.flush().await
}
fn builder<'a>() -> io::Result<Builder<'a>> {
let params: NoiseParams = NOISE_PARAMS.parse().expect("the Noise parameters should be valid");
Builder::new(params).prologue(NOISE_MAGIC.as_slice()).map_err(invalid_data)
}
fn build_state(role: Role) -> io::Result<HandshakeState> {
let private_key: [u8; DH_LEN] = rand::random();
let builder = builder()?.local_private_key(&private_key).map_err(invalid_data)?;
match role {
Role::Initiator => builder.build_initiator(),
Role::Responder => builder.build_responder(),
}
.map_err(invalid_data)
}
fn invalid_data<E: std::fmt::Display>(err: E) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, err.to_string())
}
enum SessionState {
Handshake(Box<HandshakeState>),
Transport(Box<TransportState>),
}
pub struct PendingSession<S> {
stream: S,
first_message: Vec<u8>,
}
impl<S: AsyncRead + AsyncWrite + Unpin> PendingSession<S> {
pub async fn accept(mut stream: S) -> io::Result<Self> {
let first_message = read_frame(&mut stream, MAX_HANDSHAKE_MSG_LEN).await?;
Ok(Self { stream, first_message })
}
pub fn first_payload(&self) -> io::Result<&[u8]> {
self.first_message
.get(FIRST_MESSAGE_PAYLOAD_OFFSET..)
.ok_or_else(|| invalid_data("the first Noise message is too short to carry a payload"))
}
pub fn into_session(self) -> io::Result<NoiseSession<S>> {
let Self { stream, first_message } = self;
let mut session =
NoiseSession { stream, state: SessionState::Handshake(Box::new(build_state(Role::Responder)?)) };
session.decrypt(&first_message)?;
Ok(session)
}
}
pub struct NoiseSession<S> {
stream: S,
state: SessionState,
}
impl<S: AsyncRead + AsyncWrite + Unpin> NoiseSession<S> {
pub fn new(stream: S, role: Role) -> io::Result<Self> {
Ok(Self { stream, state: SessionState::Handshake(Box::new(build_state(role)?)) })
}
pub async fn send(&mut self, payload: &[u8]) -> io::Result<()> {
let max_len = self.max_msg_len();
if payload.len() + NOISE_OVERHEAD > max_len {
return Err(invalid_data(format!("the handshake payload is too large ({} bytes)", payload.len())));
}
let mut buffer = vec![0u8; payload.len() + NOISE_OVERHEAD];
let len = match self.state {
SessionState::Handshake(ref mut state) => state.write_message(payload, &mut buffer),
SessionState::Transport(ref mut state) => state.write_message(payload, &mut buffer),
}
.map_err(invalid_data)?;
buffer.truncate(len);
write_frame(&mut self.stream, &buffer, max_len).await
}
pub async fn recv(&mut self) -> io::Result<Vec<u8>> {
let max_len = self.max_msg_len();
let message = read_frame(&mut self.stream, max_len).await?;
self.decrypt(&message)
}
fn max_msg_len(&self) -> usize {
match self.state {
SessionState::Handshake(_) => MAX_HANDSHAKE_MSG_LEN,
SessionState::Transport(_) => MAX_NOISE_MSG_LEN,
}
}
fn decrypt(&mut self, message: &[u8]) -> io::Result<Vec<u8>> {
let mut buffer = vec![0u8; message.len()];
let len = match self.state {
SessionState::Handshake(ref mut state) => state.read_message(message, &mut buffer),
SessionState::Transport(ref mut state) => state.read_message(message, &mut buffer),
}
.map_err(invalid_data)?;
buffer.truncate(len);
Ok(buffer)
}
pub fn handshake_hash(&self) -> io::Result<[u8; HANDSHAKE_HASH_LEN]> {
let SessionState::Handshake(ref state) = self.state else {
return Err(invalid_data("the Noise handshake hash is only available during the handshake"));
};
Ok(state.get_handshake_hash().try_into().expect("the Noise handshake hash should be 32 bytes long"))
}
pub fn into_transport_mode(self) -> io::Result<Self> {
let Self { stream, state } = self;
let SessionState::Handshake(handshake_state) = state else {
return Err(invalid_data("the Noise session is already in transport mode"));
};
let transport_state = handshake_state.into_transport_mode().map_err(invalid_data)?;
Ok(Self { stream, state: SessionState::Transport(Box::new(transport_state)) })
}
pub fn into_inner(self) -> S {
self.stream
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{DuplexStream, duplex};
async fn perform_xx_handshake(
payloads: [&[u8]; 3],
) -> (NoiseSession<DuplexStream>, NoiseSession<DuplexStream>, [[u8; 32]; 4]) {
let (initiator_stream, responder_stream) = duplex(1024);
let mut initiator = NoiseSession::new(initiator_stream, Role::Initiator).unwrap();
initiator.send(payloads[0]).await.unwrap();
let pending = PendingSession::accept(responder_stream).await.unwrap();
assert_eq!(pending.first_payload().unwrap(), payloads[0]);
let mut responder = pending.into_session().unwrap();
responder.send(payloads[1]).await.unwrap();
assert_eq!(initiator.recv().await.unwrap(), payloads[1]);
let (initiator_h2, responder_h2) = (initiator.handshake_hash().unwrap(), responder.handshake_hash().unwrap());
initiator.send(payloads[2]).await.unwrap();
assert_eq!(responder.recv().await.unwrap(), payloads[2]);
let (initiator_h3, responder_h3) = (initiator.handshake_hash().unwrap(), responder.handshake_hash().unwrap());
(initiator, responder, [initiator_h2, responder_h2, initiator_h3, responder_h3])
}
#[tokio::test]
async fn xx_handshake_completes_and_agrees_on_the_handshake_hash() {
let (initiator, responder, [initiator_h2, responder_h2, initiator_h3, responder_h3]) =
perform_xx_handshake([b"hint", b"responder info", b"initiator info"]).await;
assert_eq!(initiator_h2, responder_h2);
assert_eq!(initiator_h3, responder_h3);
assert_ne!(initiator_h2, initiator_h3);
let mut initiator = initiator.into_transport_mode().unwrap();
let mut responder = responder.into_transport_mode().unwrap();
responder.send(b"responder proof").await.unwrap();
assert_eq!(initiator.recv().await.unwrap(), b"responder proof");
}
#[tokio::test]
async fn handshake_hashes_differ_between_sessions() {
let (_, _, first) = perform_xx_handshake([b"", b"", b""]).await;
let (_, _, second) = perform_xx_handshake([b"", b"", b""]).await;
assert_ne!(first, second);
}
#[tokio::test]
async fn tampering_with_a_handshake_message_is_detected() {
let (mut initiator_stream, mut responder_stream) = duplex(1024);
let mut initiator = NoiseSession::new(&mut initiator_stream, Role::Initiator).unwrap();
initiator.send(b"hint").await.unwrap();
let mut responder = PendingSession::accept(&mut responder_stream).await.unwrap().into_session().unwrap();
let mut buffer = vec![0u8; MAX_NOISE_MSG_LEN];
let SessionState::Handshake(ref mut state) = responder.state else { unreachable!() };
let len = state.write_message(b"responder info", &mut buffer).unwrap();
buffer.truncate(len);
*buffer.last_mut().unwrap() ^= 1;
write_frame(&mut responder.stream, &buffer, MAX_HANDSHAKE_MSG_LEN).await.unwrap();
assert!(initiator.recv().await.is_err());
}
#[tokio::test]
async fn a_first_message_payload_is_readable_before_any_keys_are_derived() {
let (initiator_stream, responder_stream) = duplex(1024);
let mut initiator = NoiseSession::new(initiator_stream, Role::Initiator).unwrap();
initiator.send(b"a cleartext hint").await.unwrap();
let pending = PendingSession::accept(responder_stream).await.unwrap();
assert_eq!(pending.first_payload().unwrap(), b"a cleartext hint");
let mut responder = pending.into_session().unwrap();
responder.send(b"responder info").await.unwrap();
assert_eq!(initiator.recv().await.unwrap(), b"responder info");
}
#[tokio::test]
async fn tampering_with_the_cleartext_first_payload_is_detected() {
let (initiator_stream, mut initiator_wire) = duplex(1024);
let (mut responder_wire, responder_stream) = duplex(1024);
let mut initiator = NoiseSession::new(initiator_stream, Role::Initiator).unwrap();
initiator.send(b"the original hint").await.unwrap();
let mut message = read_frame(&mut initiator_wire, MAX_HANDSHAKE_MSG_LEN).await.unwrap();
message[FIRST_MESSAGE_PAYLOAD_OFFSET] ^= 1;
write_frame(&mut responder_wire, &message, MAX_HANDSHAKE_MSG_LEN).await.unwrap();
let pending = PendingSession::accept(responder_stream).await.unwrap();
assert_eq!(pending.first_payload().unwrap(), b"uhe original hint");
let mut responder = pending.into_session().unwrap();
responder.send(b"responder info").await.unwrap();
let reply = read_frame(&mut responder_wire, MAX_HANDSHAKE_MSG_LEN).await.unwrap();
write_frame(&mut initiator_wire, &reply, MAX_HANDSHAKE_MSG_LEN).await.unwrap();
assert!(initiator.recv().await.is_err());
}
#[tokio::test]
async fn a_mismatched_prologue_fails_the_handshake() {
let params: NoiseParams = NOISE_PARAMS.parse().unwrap();
let mut odd_one_out = Builder::new(params)
.prologue(b"a different marker")
.unwrap()
.local_private_key(&[0u8; DH_LEN])
.unwrap()
.build_initiator()
.unwrap();
let (mut initiator_stream, responder_stream) = duplex(1024);
let mut buffer = vec![0u8; MAX_NOISE_MSG_LEN];
let len = odd_one_out.write_message(b"hint", &mut buffer).unwrap();
buffer.truncate(len);
write_frame(&mut initiator_stream, &buffer, MAX_HANDSHAKE_MSG_LEN).await.unwrap();
let mut responder = PendingSession::accept(responder_stream).await.unwrap().into_session().unwrap();
responder.send(b"responder info").await.unwrap();
let reply = read_frame(&mut initiator_stream, MAX_HANDSHAKE_MSG_LEN).await.unwrap();
assert!(odd_one_out.read_message(&reply, &mut vec![0u8; MAX_NOISE_MSG_LEN]).is_err());
}
#[tokio::test]
async fn a_session_leaves_bytes_that_follow_a_message_on_the_stream() {
let (mut initiator_stream, responder_stream) = duplex(1024);
let mut initiator = NoiseSession::new(&mut initiator_stream, Role::Initiator).unwrap();
initiator.send(b"hint").await.unwrap();
drop(initiator);
initiator_stream.write_all(b"pipelined").await.unwrap();
let pending = PendingSession::accept(responder_stream).await.unwrap();
assert_eq!(pending.first_payload().unwrap(), b"hint");
let responder = pending.into_session().unwrap();
let mut stream = responder.into_inner();
let mut trailing = [0u8; 9];
stream.read_exact(&mut trailing).await.unwrap();
assert_eq!(&trailing, b"pipelined");
}
#[tokio::test]
async fn an_oversized_handshake_message_is_rejected_before_it_is_read() {
let (mut initiator_stream, mut responder_stream) = duplex(1024);
let length = (MAX_HANDSHAKE_MSG_LEN + 1) as u32;
initiator_stream.write_all(&length.to_le_bytes()).await.unwrap();
let error = read_frame(&mut responder_stream, MAX_HANDSHAKE_MSG_LEN).await.unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
#[tokio::test]
async fn a_truncated_first_message_is_rejected() {
let (mut initiator_stream, responder_stream) = duplex(1024);
write_frame(&mut initiator_stream, &[0u8; DH_LEN - 1], MAX_HANDSHAKE_MSG_LEN).await.unwrap();
let pending = PendingSession::accept(responder_stream).await.unwrap();
assert!(pending.first_payload().is_err());
}
#[tokio::test]
async fn an_oversized_message_length_is_rejected_before_it_is_allocated() {
let (mut initiator_stream, mut responder_stream) = duplex(1024);
initiator_stream.write_all(&(MAX_NOISE_MSG_LEN as u32 + 1).to_le_bytes()).await.unwrap();
let error = read_frame(&mut responder_stream, MAX_HANDSHAKE_MSG_LEN).await.unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
#[tokio::test]
async fn the_binding_message_is_domain_separated() {
let hash = [7u8; HANDSHAKE_HASH_LEN];
assert_ne!(binding_message(b"bft", Role::Initiator, &hash), binding_message(b"bft", Role::Responder, &hash));
assert_ne!(binding_message(b"bft", Role::Initiator, &hash), binding_message(b"router", Role::Initiator, &hash));
}
#[tokio::test]
async fn the_noise_magic_is_detected() {
let (mut initiator_stream, mut responder_stream) = duplex(1024);
write_noise_magic(&mut initiator_stream).await.unwrap();
let (protocol, leftover) = detect_handshake_protocol(&mut responder_stream).await.unwrap();
assert_eq!(protocol, HandshakeProtocol::Noise);
assert!(leftover.is_empty());
}
#[tokio::test]
async fn a_legacy_prefix_is_detected_and_returned() {
let (mut initiator_stream, mut responder_stream) = duplex(1024);
let legacy_prefix = 87u32.to_le_bytes();
initiator_stream.write_all(&legacy_prefix).await.unwrap();
let (protocol, leftover) = detect_handshake_protocol(&mut responder_stream).await.unwrap();
assert_eq!(protocol, HandshakeProtocol::Legacy);
assert_eq!(&leftover[..], &legacy_prefix[..]);
}
#[test]
fn the_noise_magic_is_an_invalid_legacy_frame_length() {
const MAX_LEGACY_HANDSHAKE_FRAME_LEN: u32 = 1024 * 1024;
assert!(u32::from_le_bytes(NOISE_MAGIC) > MAX_LEGACY_HANDSHAKE_FRAME_LEN);
}
}