use std::io;
use snow::{Builder, Keypair, TransportState};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use crate::protocol::{
flags::{ConnectionParams, MsgType},
message::{ConnectionCtx, ProtocolMessage},
proto::Proto,
};
const NOISE_PARAMS: &str = "Noise_NK_25519_AESGCM_SHA256";
fn noise_error(err: snow::Error) -> io::Error {
io::Error::new(io::ErrorKind::Other, format!("noise error: {err:?}"))
}
pub struct NoiseIdentity {
keypair: Keypair,
}
impl NoiseIdentity {
pub fn generate() -> io::Result<Self> {
let keypair = Builder::new(NOISE_PARAMS.parse().map_err(noise_error)?)
.generate_keypair()
.map_err(noise_error)?;
Ok(Self { keypair })
}
pub fn from_keypair(private: [u8; 32], public: [u8; 32]) -> Self {
Self {
keypair: Keypair {
private: private.to_vec(),
public: public.to_vec(),
},
}
}
pub fn private_key(&self) -> [u8; 32] {
let mut out = [0u8; 32];
out.copy_from_slice(&self.keypair.private);
out
}
pub fn public_key(&self) -> [u8; 32] {
let mut out = [0u8; 32];
out.copy_from_slice(&self.keypair.public);
out
}
}
async fn write_frame<STREAM>(
stream: &mut STREAM,
msg_type: MsgType,
payload: Vec<u8>,
declared_params: ConnectionParams,
) -> io::Result<()>
where
STREAM: AsyncWriteExt + Unpin,
{
let mut msg: ProtocolMessage<Vec<u8>> =
ProtocolMessage::new(ConnectionParams::NONE, msg_type, payload)?;
msg.header.reserved = declared_params.bits();
msg.write_to(stream, Proto::TCP, None).await
}
async fn read_frame<STREAM>(
stream: &mut STREAM,
expect: MsgType,
) -> io::Result<(Vec<u8>, ConnectionParams)>
where
STREAM: AsyncReadExt + Unpin,
{
let msg: ProtocolMessage<Vec<u8>> = ProtocolMessage::read_from(stream, None).await?;
if msg.msg_type() != expect {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("expected {expect:?} during handshake, got {:?}", msg.msg_type()),
));
}
let declared_params = ConnectionParams::from_bits_truncate(msg.header.reserved);
Ok((msg.payload, declared_params))
}
pub async fn perform_handshake_initiator<STREAM>(
stream: &mut STREAM,
remote_static_pubkey: &[u8; 32],
params: ConnectionParams,
) -> io::Result<(TransportState, [u8; 16], ConnectionParams)>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
{
let mut initiator = Builder::new(NOISE_PARAMS.parse().map_err(noise_error)?)
.remote_public_key(remote_static_pubkey)
.map_err(noise_error)?
.build_initiator()
.map_err(noise_error)?;
let mut buf = vec![0u8; 65535];
let len = initiator.write_message(&[], &mut buf).map_err(noise_error)?;
write_frame(stream, MsgType::Hello, buf[..len].to_vec(), params).await?;
let (ack_payload, _ack_flags) = read_frame(stream, MsgType::HelloAck).await?;
let mut scratch = vec![0u8; 65535];
initiator
.read_message(&ack_payload, &mut scratch)
.map_err(noise_error)?;
let conn_id = conn_id_from_handshake(&initiator);
let transport = initiator.into_transport_mode().map_err(noise_error)?;
Ok((transport, conn_id, params))
}
pub async fn perform_handshake_initiator_rw<R, W>(
read: &mut R,
write: &mut W,
remote_static_pubkey: &[u8; 32],
params: ConnectionParams,
) -> io::Result<(TransportState, [u8; 16], ConnectionParams)>
where
R: AsyncReadExt + Unpin,
W: AsyncWriteExt + Unpin,
{
let mut initiator = Builder::new(NOISE_PARAMS.parse().map_err(noise_error)?)
.remote_public_key(remote_static_pubkey)
.map_err(noise_error)?
.build_initiator()
.map_err(noise_error)?;
let mut buf = vec![0u8; 65535];
let len = initiator.write_message(&[], &mut buf).map_err(noise_error)?;
write_frame(write, MsgType::Hello, buf[..len].to_vec(), params).await?;
let (ack_payload, _ack_flags) = read_frame(read, MsgType::HelloAck).await?;
let mut scratch = vec![0u8; 65535];
initiator
.read_message(&ack_payload, &mut scratch)
.map_err(noise_error)?;
let conn_id = conn_id_from_handshake(&initiator);
let transport = initiator.into_transport_mode().map_err(noise_error)?;
Ok((transport, conn_id, params))
}
pub async fn perform_handshake_responder<STREAM>(
stream: &mut STREAM,
identity: &NoiseIdentity,
) -> io::Result<(TransportState, [u8; 16], ConnectionParams)>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
{
let mut responder = Builder::new(NOISE_PARAMS.parse().map_err(noise_error)?)
.local_private_key(&identity.keypair.private)
.map_err(noise_error)?
.build_responder()
.map_err(noise_error)?;
let (hello_payload, hello_params) = read_frame(stream, MsgType::Hello).await?;
let mut scratch = vec![0u8; 65535];
responder
.read_message(&hello_payload, &mut scratch)
.map_err(noise_error)?;
let mut buf = vec![0u8; 65535];
let len = responder.write_message(&[], &mut buf).map_err(noise_error)?;
write_frame(stream, MsgType::HelloAck, buf[..len].to_vec(), ConnectionParams::NONE).await?;
let conn_id = conn_id_from_handshake(&responder);
let transport = responder.into_transport_mode().map_err(noise_error)?;
Ok((transport, conn_id, hello_params))
}
pub async fn perform_handshake_responder_rw<R, W>(
read: &mut R,
write: &mut W,
identity: &NoiseIdentity,
) -> io::Result<(TransportState, [u8; 16], ConnectionParams)>
where
R: AsyncReadExt + Unpin,
W: AsyncWriteExt + Unpin,
{
let mut responder = Builder::new(NOISE_PARAMS.parse().map_err(noise_error)?)
.local_private_key(&identity.keypair.private)
.map_err(noise_error)?
.build_responder()
.map_err(noise_error)?;
let (hello_payload, hello_params) = read_frame(read, MsgType::Hello).await?;
let mut scratch = vec![0u8; 65535];
responder
.read_message(&hello_payload, &mut scratch)
.map_err(noise_error)?;
let mut buf = vec![0u8; 65535];
let len = responder.write_message(&[], &mut buf).map_err(noise_error)?;
write_frame(write, MsgType::HelloAck, buf[..len].to_vec(), ConnectionParams::NONE).await?;
let conn_id = conn_id_from_handshake(&responder);
let transport = responder.into_transport_mode().map_err(noise_error)?;
Ok((transport, conn_id, hello_params))
}
pub(crate) fn ctx_from_handshake_result(
noise: TransportState,
conn_id: [u8; 16],
params: ConnectionParams,
) -> ConnectionCtx {
ConnectionCtx {
noise,
conn_id,
next_seq: 0,
params,
insecure: params.contains(ConnectionParams::INSECURE),
}
}
pub async fn rekey_initiator<STREAM>(
stream: &mut STREAM,
old: &mut ConnectionCtx,
remote_static_pubkey: &[u8; 32],
) -> io::Result<ConnectionCtx>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
{
let signal: ProtocolMessage<()> =
ProtocolMessage::new(ConnectionParams::ENCRYPTED, MsgType::Rekey, ())?;
signal.write_to(stream, Proto::TCP, Some(old)).await?;
let (noise, conn_id, params) =
perform_handshake_initiator(stream, remote_static_pubkey, old.params).await?;
Ok(ctx_from_handshake_result(noise, conn_id, params))
}
pub async fn rekey_initiator_rw<R, W>(
read: &mut R,
write: &mut W,
old: &mut ConnectionCtx,
remote_static_pubkey: &[u8; 32],
) -> io::Result<ConnectionCtx>
where
R: AsyncReadExt + Unpin,
W: AsyncWriteExt + Unpin,
{
let signal: ProtocolMessage<()> =
ProtocolMessage::new(ConnectionParams::ENCRYPTED, MsgType::Rekey, ())?;
signal.write_to(write, Proto::TCP, Some(old)).await?;
let (noise, conn_id, params) =
perform_handshake_initiator_rw(read, write, remote_static_pubkey, old.params).await?;
Ok(ctx_from_handshake_result(noise, conn_id, params))
}
pub async fn rekey_responder<STREAM>(
stream: &mut STREAM,
identity: &NoiseIdentity,
) -> io::Result<ConnectionCtx>
where
STREAM: AsyncReadExt + AsyncWriteExt + Unpin,
{
let (noise, conn_id, params) = perform_handshake_responder(stream, identity).await?;
Ok(ctx_from_handshake_result(noise, conn_id, params))
}
pub async fn rekey_responder_rw<R, W>(
read: &mut R,
write: &mut W,
identity: &NoiseIdentity,
) -> io::Result<ConnectionCtx>
where
R: AsyncReadExt + Unpin,
W: AsyncWriteExt + Unpin,
{
let (noise, conn_id, params) = perform_handshake_responder_rw(read, write, identity).await?;
Ok(ctx_from_handshake_result(noise, conn_id, params))
}
fn conn_id_from_handshake(state: &snow::HandshakeState) -> [u8; 16] {
let hash = state.get_handshake_hash();
let mut conn_id = [0u8; 16];
conn_id.copy_from_slice(&hash[..16]);
conn_id
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::io_helpers::read_until;
use crate::protocol::header::EOL;
#[tokio::test]
async fn handshake_agrees_on_connection_id() {
let identity = NoiseIdentity::generate().unwrap();
let remote_pub = identity.public_key();
let (mut client, mut server) = tokio::io::duplex(4096);
let (initiator_result, responder_result) = tokio::join!(
perform_handshake_initiator(&mut client, &remote_pub, ConnectionParams::INSECURE),
perform_handshake_responder(&mut server, &identity),
);
let (initiator_transport, initiator_conn_id, initiator_params) = initiator_result.unwrap();
let (_responder_transport, responder_conn_id, responder_params) = responder_result.unwrap();
assert_eq!(initiator_conn_id, responder_conn_id);
assert_eq!(initiator_params, responder_params);
assert!(responder_params.contains(ConnectionParams::INSECURE));
assert!(initiator_transport.is_initiator());
}
#[tokio::test]
async fn rekey_produces_a_fresh_connection() {
let identity = NoiseIdentity::generate().unwrap();
let remote_pub = identity.public_key();
let (mut client, mut server) = tokio::io::duplex(4096);
let client_fut = async {
let (noise, conn_id, params) =
perform_handshake_initiator(&mut client, &remote_pub, ConnectionParams::ENCRYPTED)
.await
.unwrap();
let mut client_ctx = ctx_from_handshake_result(noise, conn_id, params);
let old_conn_id = client_ctx.conn_id;
let new_ctx = rekey_initiator(&mut client, &mut client_ctx, &remote_pub)
.await
.unwrap();
(old_conn_id, new_ctx)
};
let server_fut = async {
let (noise, conn_id, params) = perform_handshake_responder(&mut server, &identity)
.await
.unwrap();
let mut server_ctx = ctx_from_handshake_result(noise, conn_id, params);
let old_conn_id = server_ctx.conn_id;
let mut buffer = read_until(&mut server, EOL.to_vec()).await.unwrap();
if let Some(pos) = buffer.windows(EOL.len()).rposition(|w| w == EOL) {
buffer.truncate(pos);
}
let signal: ProtocolMessage<()> =
ProtocolMessage::from_bytes(&buffer, Some(&mut server_ctx)).unwrap();
assert_eq!(signal.msg_type(), MsgType::Rekey);
let new_ctx = rekey_responder(&mut server, &identity).await.unwrap();
(old_conn_id, new_ctx)
};
let ((old_client_conn_id, new_client_ctx), (old_server_conn_id, mut new_server_ctx)) =
tokio::join!(client_fut, server_fut);
assert_eq!(old_client_conn_id, old_server_conn_id);
assert_ne!(new_client_ctx.conn_id, old_client_conn_id);
assert_eq!(new_client_ctx.conn_id, new_server_ctx.conn_id);
assert_eq!(new_client_ctx.params, ConnectionParams::ENCRYPTED);
let mut client_ctx = new_client_ctx;
let mut data_msg =
ProtocolMessage::new(ConnectionParams::ENCRYPTED, MsgType::Data, b"post-rekey".to_vec())
.unwrap();
let bytes = data_msg.to_bytes(Some(&mut client_ctx)).unwrap();
let parsed: ProtocolMessage<Vec<u8>> =
ProtocolMessage::from_bytes(&bytes, Some(&mut new_server_ctx)).unwrap();
assert_eq!(parsed.payload, b"post-rekey".to_vec());
}
}