use std::{
cmp,
convert::TryInto,
io,
pin::Pin,
task::{Context, Poll},
time::Duration,
};
use futures::ready;
use log::*;
use snow::{HandshakeState, TransportState, error::StateProblem};
use tari_utilities::ByteArray;
use tokio::{
io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf},
time,
};
use crate::types::CommsPublicKey;
const LOG_TARGET: &str = "comms::noise::socket";
const MAX_PAYLOAD_LENGTH: usize = u16::MAX as usize;
pub(crate) const MAX_WRITE_BUFFER_LENGTH: usize = u16::MAX as usize - 16;
const FRAME_LEN_PREFIX_LENGTH: usize = 2;
const READ_AHEAD_LENGTH: usize = 128 * 1024;
struct NoiseBuffers {
read_ahead: Box<[u8]>,
read_start: usize,
read_end: usize,
read_decrypted: [u8; MAX_PAYLOAD_LENGTH],
write_decrypted: [u8; MAX_WRITE_BUFFER_LENGTH],
write_encrypted: [u8; FRAME_LEN_PREFIX_LENGTH + MAX_PAYLOAD_LENGTH],
}
impl NoiseBuffers {
fn new() -> Self {
Self {
read_ahead: vec![0; READ_AHEAD_LENGTH].into_boxed_slice(),
read_start: 0,
read_end: 0,
read_decrypted: [0; MAX_PAYLOAD_LENGTH],
write_decrypted: [0; MAX_WRITE_BUFFER_LENGTH],
write_encrypted: [0; FRAME_LEN_PREFIX_LENGTH + MAX_PAYLOAD_LENGTH],
}
}
fn read_available(&self) -> usize {
self.read_end.saturating_sub(self.read_start)
}
fn read_unconsumed(&self) -> &[u8] {
self.read_ahead.get(self.read_start..self.read_end).unwrap_or(&[])
}
fn read_consume(&mut self, n: usize) {
self.read_start = cmp::min(self.read_start.saturating_add(n), self.read_end);
if self.read_start == self.read_end {
self.read_start = 0;
self.read_end = 0;
}
}
}
impl ::std::fmt::Debug for NoiseBuffers {
fn fmt(&self, f: &mut ::std::fmt::Formatter) -> ::std::fmt::Result {
f.debug_struct("NoiseBuffers").finish()
}
}
#[derive(Debug)]
enum ReadState {
Init,
ReadFrameLen,
ReadFrame { frame_len: u16 },
CopyDecryptedFrame { decrypted_len: usize, offset: usize },
Eof(Result<(), ()>),
DecryptionError(snow::Error),
}
#[derive(Debug)]
enum WriteState {
Init,
BufferData { offset: usize },
WriteEncryptedFrame { frame_len: u16, offset: usize },
Flush,
Eof,
EncryptionError(snow::Error),
}
#[derive(Debug)]
pub struct NoiseSocket<TSocket> {
socket: TSocket,
state: NoiseState,
buffers: Box<NoiseBuffers>,
read_state: ReadState,
write_state: WriteState,
}
impl<TSocket> NoiseSocket<TSocket> {
fn new(socket: TSocket, session: NoiseState) -> Self {
Self {
socket,
state: session,
buffers: Box::new(NoiseBuffers::new()),
read_state: ReadState::Init,
write_state: WriteState::Init,
}
}
pub fn get_remote_static(&self) -> Option<&[u8]> {
self.state.get_remote_static()
}
pub fn get_remote_public_key(&self) -> Option<CommsPublicKey> {
self.get_remote_static()
.and_then(|s| CommsPublicKey::from_canonical_bytes(s).ok())
}
fn is_eof(&self) -> bool {
matches!(self.read_state, ReadState::Eof(_))
}
}
fn poll_write_all<TSocket>(
context: &mut Context,
mut socket: Pin<&mut TSocket>,
buf: &[u8],
offset: &mut usize,
) -> Poll<io::Result<()>>
where
TSocket: AsyncWrite,
{
loop {
let bytes = match buf.get(*offset..) {
Some(bytes) => bytes,
None => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidInput,
"Offset exceeds buffer length",
)));
},
};
let n = ready!(socket.as_mut().poll_write(context, bytes))?;
trace!(
target: LOG_TARGET,
"poll_write_all: wrote {}/{} bytes",
offset.saturating_add(n),
buf.len()
);
if n == 0 {
return Poll::Ready(Err(io::ErrorKind::WriteZero.into()));
}
*offset = offset.saturating_add(n);
assert!(*offset <= buf.len());
if *offset == buf.len() {
return Poll::Ready(Ok(()));
}
}
}
fn poll_fill_read_ahead<TSocket>(
context: &mut Context,
socket: Pin<&mut TSocket>,
buffers: &mut NoiseBuffers,
needed: usize,
) -> Poll<io::Result<usize>>
where
TSocket: AsyncRead,
{
if buffers.read_start.saturating_add(needed) > buffers.read_ahead.len() {
buffers.read_ahead.copy_within(buffers.read_start..buffers.read_end, 0);
buffers.read_end = buffers.read_available();
buffers.read_start = 0;
}
let free = match buffers.read_ahead.get_mut(buffers.read_end..) {
Some(free) if !free.is_empty() => free,
_ => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidInput,
"read-ahead buffer is full",
)));
},
};
let mut read_buf = ReadBuf::new(free);
ready!(socket.poll_read(context, &mut read_buf))?;
let n = read_buf.filled().len();
buffers.read_end = buffers.read_end.saturating_add(n);
trace!(
target: LOG_TARGET,
"poll_fill_read_ahead: read {} bytes, {} bytes buffered",
n,
buffers.read_available()
);
Poll::Ready(Ok(n))
}
impl<TSocket> NoiseSocket<TSocket>
where TSocket: AsyncRead + Unpin
{
#[allow(clippy::too_many_lines)]
fn poll_read(&mut self, context: &mut Context, buf: &mut [u8]) -> Poll<io::Result<usize>> {
loop {
trace!(target: LOG_TARGET, "NoiseSocket ReadState::{:?}", self.read_state);
match self.read_state {
ReadState::Init => {
self.read_state = ReadState::ReadFrameLen;
},
ReadState::ReadFrameLen => {
let len_bytes = self.buffers.read_unconsumed().get(..2).and_then(|b| b.try_into().ok());
if let Some(len_bytes) = len_bytes {
let coop = ready!(tokio::task::coop::poll_proceed(context));
coop.made_progress();
let frame_len = u16::from_be_bytes(len_bytes);
self.buffers.read_consume(2);
if frame_len == 0 {
self.read_state = ReadState::Init;
} else {
self.read_state = ReadState::ReadFrame { frame_len };
}
continue;
}
let n = ready!(poll_fill_read_ahead(
context,
Pin::new(&mut self.socket),
&mut self.buffers,
2
))?;
if n == 0 {
if self.buffers.read_available() == 0 {
self.read_state = ReadState::Eof(Ok(()));
} else {
self.read_state = ReadState::Eof(Err(()));
return Poll::Ready(Err(io::ErrorKind::UnexpectedEof.into()));
}
}
},
ReadState::ReadFrame { frame_len } => {
let frame_len = usize::from(frame_len);
if self.buffers.read_available() >= frame_len {
let coop = ready!(tokio::task::coop::poll_proceed(context));
coop.made_progress();
let NoiseBuffers {
read_ahead,
read_start,
read_decrypted,
..
} = &mut *self.buffers;
let frame = read_ahead
.get(*read_start..read_start.saturating_add(frame_len))
.expect("this is checked");
let result = self.state.read_message(frame, read_decrypted);
self.buffers.read_consume(frame_len);
match result {
Ok(decrypted_len) => {
self.read_state = ReadState::CopyDecryptedFrame {
decrypted_len,
offset: 0,
};
},
Err(e) => {
warn!(target: LOG_TARGET, "Decryption Error: {e}");
self.read_state = ReadState::DecryptionError(e);
},
}
continue;
}
let n = ready!(poll_fill_read_ahead(
context,
Pin::new(&mut self.socket),
&mut self.buffers,
frame_len
))?;
if n == 0 {
self.read_state = ReadState::Eof(Err(()));
return Poll::Ready(Err(io::ErrorKind::UnexpectedEof.into()));
}
},
ReadState::CopyDecryptedFrame {
decrypted_len,
ref mut offset,
} => {
let num_bytes_to_copy = cmp::min(decrypted_len.saturating_sub(*offset), buf.len());
let copy_end = offset.saturating_add(num_bytes_to_copy);
let bytes_to_copy = match self.buffers.read_decrypted.get(*offset..copy_end) {
Some(bytes) => bytes,
None => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidInput,
"Offset exceeds buffer length",
)));
},
};
buf.get_mut(..num_bytes_to_copy)
.expect("this is checked")
.copy_from_slice(bytes_to_copy);
trace!(
target: LOG_TARGET,
"CopyDecryptedFrame: copied {}/{} bytes",
copy_end,
decrypted_len
);
*offset = copy_end;
if *offset == decrypted_len {
self.read_state = ReadState::Init;
}
return Poll::Ready(Ok(num_bytes_to_copy));
},
ReadState::Eof(Ok(())) => return Poll::Ready(Ok(0)),
ReadState::Eof(Err(())) => return Poll::Ready(Err(io::ErrorKind::UnexpectedEof.into())),
ReadState::DecryptionError(ref e) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("DecryptionError: {e}"),
)));
},
}
}
}
}
impl<TSocket> AsyncRead for NoiseSocket<TSocket>
where TSocket: AsyncRead + Unpin
{
fn poll_read(self: Pin<&mut Self>, context: &mut Context, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
let slice = buf.initialize_unfilled();
let n = futures::ready!(self.get_mut().poll_read(context, slice))?;
buf.advance(n);
Poll::Ready(Ok(()))
}
}
impl<TSocket> NoiseSocket<TSocket>
where TSocket: AsyncWrite + Unpin
{
#[allow(clippy::too_many_lines)]
fn poll_write_or_flush(&mut self, context: &mut Context, buf: Option<&[u8]>) -> Poll<io::Result<Option<usize>>> {
loop {
trace!(
target: LOG_TARGET,
"NoiseSocket {} WriteState::{:?}",
if buf.is_some() { "poll_write" } else { "poll_flush" },
self.write_state,
);
match self.write_state {
WriteState::Init => {
if buf.is_some() {
self.write_state = WriteState::BufferData { offset: 0 };
} else {
return Poll::Ready(Ok(None));
}
},
WriteState::BufferData { ref mut offset } => {
let bytes_buffered = if let Some(buf) = buf {
let num_bytes_to_copy =
::std::cmp::min(MAX_WRITE_BUFFER_LENGTH.saturating_sub(*offset), buf.len());
let bytes = match buf.get(..num_bytes_to_copy) {
Some(bytes) => bytes,
None => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidInput,
"frame length exceeds buffer length",
)));
},
};
self.buffers
.write_decrypted
.get_mut(*offset..offset.saturating_add(num_bytes_to_copy))
.expect("this is checked")
.copy_from_slice(bytes);
trace!(
target: LOG_TARGET,
"BufferData: buffered {}/{} bytes",
num_bytes_to_copy,
buf.len()
);
*offset = offset.saturating_add(num_bytes_to_copy);
Some(num_bytes_to_copy)
} else {
None
};
if buf.is_none() || *offset == MAX_WRITE_BUFFER_LENGTH {
let bytes = match self.buffers.write_decrypted.get(..*offset) {
Some(bytes) => bytes,
None => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidInput,
"frame length exceeds buffer length",
)));
},
};
let (len_prefix, frame) = self.buffers.write_encrypted.split_at_mut(FRAME_LEN_PREFIX_LENGTH);
match self.state.write_message(bytes, frame) {
Ok(encrypted_len) => {
let frame_len: u16 = encrypted_len
.try_into()
.map_err(|_| io::Error::other("offset should be able to fit in u16"))?;
len_prefix.copy_from_slice(&frame_len.to_be_bytes());
self.write_state = WriteState::WriteEncryptedFrame { frame_len, offset: 0 };
},
Err(e) => {
warn!(target: LOG_TARGET, "Encryption Error: {e}");
let err = io::Error::new(io::ErrorKind::InvalidData, format!("EncryptionError: {e}"));
self.write_state = WriteState::EncryptionError(e);
return Poll::Ready(Err(err));
},
}
}
if let Some(bytes_buffered) = bytes_buffered {
return Poll::Ready(Ok(Some(bytes_buffered)));
}
},
WriteState::WriteEncryptedFrame {
frame_len,
ref mut offset,
} => {
let bytes = match self
.buffers
.write_encrypted
.get(..FRAME_LEN_PREFIX_LENGTH.saturating_add(frame_len as usize))
{
Some(bytes) => bytes,
None => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidInput,
"frame length exceeds buffer length",
)));
},
};
match ready!(poll_write_all(context, Pin::new(&mut self.socket), bytes, offset)) {
Ok(()) => {
self.write_state = WriteState::Flush;
},
Err(e) => {
if e.kind() == io::ErrorKind::WriteZero {
self.write_state = WriteState::Eof;
}
return Poll::Ready(Err(e));
},
}
},
WriteState::Flush => {
ready!(Pin::new(&mut self.socket).poll_flush(context))?;
self.write_state = WriteState::Init;
},
WriteState::Eof => return Poll::Ready(Err(io::ErrorKind::WriteZero.into())),
WriteState::EncryptionError(ref e) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("EncryptionError: {e}"),
)));
},
}
}
}
fn poll_write(&mut self, context: &mut Context, buf: &[u8]) -> Poll<io::Result<usize>> {
if let Some(bytes_written) = ready!(self.poll_write_or_flush(context, Some(buf)))? {
Poll::Ready(Ok(bytes_written))
} else {
unreachable!();
}
}
fn poll_flush(&mut self, context: &mut Context) -> Poll<io::Result<()>> {
if ready!(self.poll_write_or_flush(context, None))?.is_none() {
Poll::Ready(Ok(()))
} else {
unreachable!();
}
}
}
impl<TSocket> AsyncWrite for NoiseSocket<TSocket>
where TSocket: AsyncWrite + Unpin
{
fn poll_write(self: Pin<&mut Self>, cx: &mut Context, buf: &[u8]) -> Poll<io::Result<usize>> {
self.get_mut().poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
self.get_mut().poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
Pin::new(&mut self.socket).poll_shutdown(cx)
}
}
pub struct Handshake<TSocket> {
socket: NoiseSocket<TSocket>,
recv_timeout: Duration,
}
impl<TSocket> Handshake<TSocket> {
pub fn new(socket: TSocket, state: HandshakeState, recv_timeout: Duration) -> Self {
Self {
socket: NoiseSocket::new(socket, state.into()),
recv_timeout,
}
}
}
impl<TSocket> Handshake<TSocket>
where TSocket: AsyncRead + AsyncWrite + Unpin
{
pub async fn perform_handshake(mut self) -> io::Result<NoiseSocket<TSocket>> {
match self.handshake_1_5rtt().await {
Ok(_) => self.build(),
Err(err) => {
info!(
target: LOG_TARGET,
"Noise handshake failed because '{err:?}'. Closing socket."
);
self.socket.shutdown().await?;
Err(err)
},
}
}
async fn handshake_1_5rtt(&mut self) -> io::Result<()> {
if self.socket.state.is_initiator() {
self.send().await?;
self.flush().await?;
self.receive().await?;
self.send().await?;
self.flush().await?;
} else {
self.receive().await?;
self.send().await?;
self.flush().await?;
self.receive().await?;
}
Ok(())
}
async fn send(&mut self) -> io::Result<usize> {
self.socket.write(&[]).await
}
async fn flush(&mut self) -> io::Result<()> {
self.socket.flush().await
}
async fn receive(&mut self) -> io::Result<usize> {
let num_bytes = time::timeout(self.recv_timeout, self.socket.read(&mut []))
.await
.map_err(|_| io::Error::from(io::ErrorKind::TimedOut))??;
if self.socket.is_eof() {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"peer closed the connection during the noise handshake",
));
}
Ok(num_bytes)
}
fn build(self) -> io::Result<NoiseSocket<TSocket>> {
let transport_state = self
.socket
.state
.into_transport_mode()
.map_err(|err| io::Error::other(format!("Invalid snow state: {err}")))?;
Ok(NoiseSocket {
state: transport_state,
..self.socket
})
}
}
#[derive(Debug)]
enum NoiseState {
HandshakeState(Box<HandshakeState>),
TransportState(Box<TransportState>),
}
macro_rules! proxy_state_method {
(pub fn $name:ident(&mut self$(,)? $($arg_name:ident : $arg_type:ty),*) -> $ret:ty) => {
pub fn $name(&mut self, $($arg_name:$arg_type),*) -> $ret {
match self {
NoiseState::HandshakeState(state) => state.$name($($arg_name),*),
NoiseState::TransportState(state) => state.$name($($arg_name),*),
}
}
};
(pub fn $name:ident(&self$(,)? $($arg_name:ident : $arg_type:ty),*) -> $ret:ty) => {
pub fn $name(&self, $($arg_name:$arg_type),*) -> $ret {
match self {
NoiseState::HandshakeState(state) => state.$name($($arg_name),*),
NoiseState::TransportState(state) => state.$name($($arg_name),*),
}
}
}
}
impl NoiseState {
proxy_state_method!(pub fn write_message(&mut self, message: &[u8], payload: &mut [u8]) -> Result<usize, snow::Error>);
proxy_state_method!(pub fn is_initiator(&self) -> bool);
proxy_state_method!(pub fn read_message(&mut self, message: &[u8], payload: &mut [u8]) -> Result<usize, snow::Error>);
proxy_state_method!(pub fn get_remote_static(&self) -> Option<&[u8]>);
pub fn into_transport_mode(self) -> Result<Self, snow::Error> {
match self {
NoiseState::HandshakeState(state) => Ok(NoiseState::TransportState(Box::new(state.into_transport_mode()?))),
_ => Err(snow::Error::State(StateProblem::HandshakeAlreadyFinished)),
}
}
}
impl From<HandshakeState> for NoiseState {
fn from(state: HandshakeState) -> Self {
NoiseState::HandshakeState(Box::new(state))
}
}
impl From<TransportState> for NoiseState {
fn from(state: TransportState) -> Self {
NoiseState::TransportState(Box::new(state))
}
}
#[cfg(test)]
mod test {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use futures::future::join;
use snow::{Builder, Error, Keypair, params::NoiseParams};
use tokio::sync::mpsc;
use super::*;
use crate::{memsocket::MemorySocket, noise::config::NOISE_PARAMETERS};
async fn build_test_connection()
-> Result<((Keypair, Handshake<MemorySocket>), (Keypair, Handshake<MemorySocket>)), Error> {
let parameters: NoiseParams = NOISE_PARAMETERS.parse().expect("Invalid protocol name");
let dialer_keypair = Builder::new(parameters.clone()).generate_keypair()?;
let listener_keypair = Builder::new(parameters.clone()).generate_keypair()?;
let dialer_session = Builder::new(parameters.clone())
.local_private_key(&dialer_keypair.private)
.build_initiator()?;
let listener_session = Builder::new(parameters)
.local_private_key(&listener_keypair.private)
.build_responder()?;
let (dialer_socket, listener_socket) = MemorySocket::new_pair();
let (dialer, listener) = (
NoiseSocket::new(dialer_socket, dialer_session.into()),
NoiseSocket::new(listener_socket, listener_session.into()),
);
Ok((
(dialer_keypair, Handshake {
socket: dialer,
recv_timeout: Duration::from_secs(1),
}),
(listener_keypair, Handshake {
socket: listener,
recv_timeout: Duration::from_secs(1),
}),
))
}
async fn perform_handshake(
dialer: Handshake<MemorySocket>,
listener: Handshake<MemorySocket>,
) -> io::Result<(NoiseSocket<MemorySocket>, NoiseSocket<MemorySocket>)> {
let (dialer_result, listener_result) = join(dialer.perform_handshake(), listener.perform_handshake()).await;
Ok((dialer_result?, listener_result?))
}
#[tokio::test]
async fn test_handshake() {
let ((dialer_keypair, dialer), (listener_keypair, listener)) = build_test_connection().await.unwrap();
let (dialer_socket, listener_socket) = perform_handshake(dialer, listener).await.unwrap();
assert_eq!(
dialer_socket.get_remote_static(),
Some(listener_keypair.public.as_ref())
);
assert_eq!(
listener_socket.get_remote_static(),
Some(dialer_keypair.public.as_ref())
);
}
#[tokio::test]
async fn handshake_reports_eof_when_peer_hangs_up() {
let ((_dialer_keypair, dialer), (_listener_keypair, mut listener)) = build_test_connection().await.unwrap();
let listener_task = tokio::spawn(async move {
listener.receive().await.unwrap();
drop(listener);
});
let err = dialer.perform_handshake().await.unwrap_err();
listener_task.await.unwrap();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof, "unexpected error: {err}");
}
#[tokio::test]
async fn simple_test() -> io::Result<()> {
let ((_dialer_keypair, dialer), (_listener_keypair, listener)) = build_test_connection().await.unwrap();
let (mut dialer_socket, mut listener_socket) = perform_handshake(dialer, listener).await?;
dialer_socket.write_all(b"stormlight").await?;
dialer_socket.write_all(b" ").await?;
dialer_socket.write_all(b"archive").await?;
dialer_socket.flush().await?;
dialer_socket.shutdown().await?;
let mut buf = Vec::new();
listener_socket.read_to_end(&mut buf).await?;
assert_eq!(buf, b"stormlight archive");
Ok(())
}
#[tokio::test]
async fn interleaved_writes() -> io::Result<()> {
let ((_dialer_keypair, dialer), (_listener_keypair, listener)) = build_test_connection().await.unwrap();
let (mut a, mut b) = perform_handshake(dialer, listener).await?;
a.write_all(b"The Name of the Wind").await?;
a.flush().await?;
a.write_all(b"The Wise Man's Fear").await?;
a.flush().await?;
b.write_all(b"The Doors of Stone").await?;
b.flush().await?;
let mut buf = [0; 20];
b.read_exact(&mut buf).await?;
assert_eq!(&buf, b"The Name of the Wind");
let mut buf = [0; 19];
b.read_exact(&mut buf).await?;
assert_eq!(&buf, b"The Wise Man's Fear");
let mut buf = [0; 18];
a.read_exact(&mut buf).await?;
assert_eq!(&buf, b"The Doors of Stone");
Ok(())
}
#[tokio::test]
async fn u16_max_writes() -> io::Result<()> {
let ((_dialer_keypair, dialer), (_listener_keypair, listener)) = build_test_connection().await.unwrap();
let (mut a, mut b) = perform_handshake(dialer, listener).await?;
let buf_send = &[1; MAX_PAYLOAD_LENGTH + 1];
a.write_all(buf_send).await?;
a.flush().await?;
let mut buf_receive = vec![0; MAX_PAYLOAD_LENGTH + 1];
b.read_exact(&mut buf_receive).await?;
assert_eq!(&buf_receive[..], &buf_send[..]);
Ok(())
}
#[tokio::test]
async fn larger_writes() -> io::Result<()> {
let ((_dialer_keypair, dialer), (_listener_keypair, listener)) = build_test_connection().await.unwrap();
let (mut a, mut b) = perform_handshake(dialer, listener).await?;
let buf_send = &[1; MAX_PAYLOAD_LENGTH * 2 + 1024];
a.write_all(buf_send).await?;
a.flush().await?;
let mut buf_receive = vec![0; MAX_PAYLOAD_LENGTH * 2 + 1024];
b.read_exact(&mut buf_receive).await?;
assert_eq!(&buf_receive[..], &buf_send[..]);
Ok(())
}
struct WriteCountingSocket {
inner: MemorySocket,
num_writes: Arc<AtomicUsize>,
}
impl AsyncRead for WriteCountingSocket {
fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl AsyncWrite for WriteCountingSocket {
fn poll_write(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
self.num_writes.fetch_add(1, Ordering::SeqCst);
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
}
#[tokio::test]
async fn one_write_per_frame() -> io::Result<()> {
let parameters: NoiseParams = NOISE_PARAMETERS.parse().expect("Invalid protocol name");
let dialer_keypair = Builder::new(parameters.clone()).generate_keypair().unwrap();
let listener_keypair = Builder::new(parameters.clone()).generate_keypair().unwrap();
let dialer_session = Builder::new(parameters.clone())
.local_private_key(&dialer_keypair.private)
.build_initiator()
.unwrap();
let listener_session = Builder::new(parameters)
.local_private_key(&listener_keypair.private)
.build_responder()
.unwrap();
let (dialer_socket, listener_socket) = MemorySocket::new_pair();
let num_writes = Arc::new(AtomicUsize::new(0));
let dialer_socket = WriteCountingSocket {
inner: dialer_socket,
num_writes: num_writes.clone(),
};
let dialer = Handshake::new(dialer_socket, dialer_session, Duration::from_secs(1));
let listener = Handshake::new(listener_socket, listener_session, Duration::from_secs(1));
let (a, b) = join(dialer.perform_handshake(), listener.perform_handshake()).await;
let (mut a, mut b) = (a?, b?);
assert_eq!(num_writes.swap(0, Ordering::SeqCst), 2);
a.write_all(b"one frame").await?;
a.flush().await?;
assert_eq!(num_writes.swap(0, Ordering::SeqCst), 1);
let buf_send = vec![7u8; MAX_WRITE_BUFFER_LENGTH * 2 + 100];
a.write_all(&buf_send).await?;
a.flush().await?;
assert_eq!(num_writes.swap(0, Ordering::SeqCst), 3);
let mut buf_receive = [0u8; 9];
b.read_exact(&mut buf_receive).await?;
assert_eq!(&buf_receive, b"one frame");
let mut buf_receive = vec![0u8; buf_send.len()];
b.read_exact(&mut buf_receive).await?;
assert_eq!(buf_receive, buf_send);
Ok(())
}
struct ShortWriteSocket {
inner: MemorySocket,
num_calls: usize,
}
impl AsyncRead for ShortWriteSocket {
fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl AsyncWrite for ShortWriteSocket {
fn poll_write(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
self.num_calls = self.num_calls.wrapping_add(1);
let max_len = match self.num_calls % 3 {
0 => {
cx.waker().wake_by_ref();
return Poll::Pending;
},
1 => 1,
_ => 7,
};
let len = cmp::min(buf.len(), max_len);
Pin::new(&mut self.inner).poll_write(cx, buf.get(..len).unwrap_or(buf))
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
}
#[tokio::test]
async fn partial_writes_resume_mid_frame() -> io::Result<()> {
let parameters: NoiseParams = NOISE_PARAMETERS.parse().expect("Invalid protocol name");
let dialer_keypair = Builder::new(parameters.clone()).generate_keypair().unwrap();
let listener_keypair = Builder::new(parameters.clone()).generate_keypair().unwrap();
let dialer_session = Builder::new(parameters.clone())
.local_private_key(&dialer_keypair.private)
.build_initiator()
.unwrap();
let listener_session = Builder::new(parameters)
.local_private_key(&listener_keypair.private)
.build_responder()
.unwrap();
let (dialer_socket, listener_socket) = MemorySocket::new_pair();
let dialer_socket = ShortWriteSocket {
inner: dialer_socket,
num_calls: 0,
};
let dialer = Handshake::new(dialer_socket, dialer_session, Duration::from_secs(1));
let listener = Handshake::new(listener_socket, listener_session, Duration::from_secs(1));
let (a, b) = join(dialer.perform_handshake(), listener.perform_handshake()).await;
let (mut a, mut b) = (a?, b?);
let buf_send = (0..MAX_WRITE_BUFFER_LENGTH * 2 + 100)
.map(|i| u8::try_from(i % 251).unwrap())
.collect::<Vec<_>>();
a.write_all(&buf_send).await?;
a.flush().await?;
let mut buf_receive = vec![0u8; buf_send.len()];
b.read_exact(&mut buf_receive).await?;
assert_eq!(buf_receive, buf_send);
Ok(())
}
#[tokio::test]
async fn unexpected_eof() -> io::Result<()> {
let ((_dialer_keypair, dialer), (_listener_keypair, listener)) = build_test_connection().await.unwrap();
let (mut a, mut b) = perform_handshake(dialer, listener).await?;
let buf_send = &[1; MAX_PAYLOAD_LENGTH];
a.write_all(buf_send).await?;
a.flush().await?;
a.socket.shutdown().await.unwrap();
drop(a);
let mut buf_receive = vec![0; MAX_PAYLOAD_LENGTH];
b.read_exact(&mut buf_receive).await.unwrap();
assert_eq!(&buf_receive[..], &buf_send[..]);
let err = b.read_exact(&mut buf_receive).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
Ok(())
}
struct MockSocket {
inbound: mpsc::UnboundedReceiver<Vec<u8>>,
pending: Vec<u8>,
outbound: mpsc::UnboundedSender<Vec<u8>>,
reads: Arc<AtomicUsize>,
}
struct MockPeer {
inbound: mpsc::UnboundedSender<Vec<u8>>,
outbound: mpsc::UnboundedReceiver<Vec<u8>>,
reads: Arc<AtomicUsize>,
}
impl MockSocket {
fn new() -> (Self, MockPeer) {
let (in_tx, in_rx) = mpsc::unbounded_channel();
let (out_tx, out_rx) = mpsc::unbounded_channel();
let reads = Arc::new(AtomicUsize::new(0));
let socket = Self {
inbound: in_rx,
pending: Vec::new(),
outbound: out_tx,
reads: reads.clone(),
};
let peer = MockPeer {
inbound: in_tx,
outbound: out_rx,
reads,
};
(socket, peer)
}
}
impl AsyncRead for MockSocket {
fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
if self.pending.is_empty() {
match ready!(self.inbound.poll_recv(cx)) {
Some(chunk) => self.pending = chunk,
None => return Poll::Ready(Ok(())),
}
}
let n = cmp::min(self.pending.len(), buf.remaining());
let chunk: Vec<u8> = self.pending.drain(..n).collect();
buf.put_slice(&chunk);
self.reads.fetch_add(1, Ordering::SeqCst);
Poll::Ready(Ok(()))
}
}
impl AsyncWrite for MockSocket {
fn poll_write(self: Pin<&mut Self>, _cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
let _ignore = self.outbound.send(buf.to_vec());
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
fn handshake_states() -> (HandshakeState, HandshakeState) {
let parameters: NoiseParams = NOISE_PARAMETERS.parse().unwrap();
let initiator_keypair = Builder::new(parameters.clone()).generate_keypair().unwrap();
let responder_keypair = Builder::new(parameters.clone()).generate_keypair().unwrap();
let initiator = Builder::new(parameters.clone())
.local_private_key(&initiator_keypair.private)
.build_initiator()
.unwrap();
let responder = Builder::new(parameters)
.local_private_key(&responder_keypair.private)
.build_responder()
.unwrap();
(initiator, responder)
}
fn wire_frame(state: &mut NoiseState, message: &[u8]) -> Vec<u8> {
let mut encrypted = vec![0u8; MAX_PAYLOAD_LENGTH];
let len = state.write_message(message, &mut encrypted).unwrap();
encrypted.truncate(len);
let mut frame = u16::try_from(len).unwrap().to_be_bytes().to_vec();
frame.extend(encrypted);
frame
}
fn transport_states() -> (NoiseState, NoiseState) {
let (initiator, responder) = handshake_states();
let mut initiator = NoiseState::from(initiator);
let mut responder = NoiseState::from(responder);
let mut payload = vec![0u8; MAX_PAYLOAD_LENGTH];
for (sender, receiver) in [(0, 1), (1, 0), (0, 1)] {
let mut states = [&mut initiator, &mut responder];
let frame = wire_frame(states.get_mut(sender).unwrap(), &[]);
let message = frame.get(2..).unwrap();
states
.get_mut(receiver)
.unwrap()
.read_message(message, &mut payload)
.unwrap();
}
(
initiator.into_transport_mode().unwrap(),
responder.into_transport_mode().unwrap(),
)
}
#[tokio::test]
async fn several_frames_in_one_read_are_decoded_in_order() {
let (mut sender, receiver) = transport_states();
let (
mock,
MockPeer {
inbound: in_tx,
outbound: _out_rx,
reads,
},
) = MockSocket::new();
let mut socket = NoiseSocket::new(mock, receiver);
let mut chunk = Vec::new();
for message in [&b"first "[..], b"second ", b"third"] {
chunk.extend(wire_frame(&mut sender, message));
}
in_tx.send(chunk).unwrap();
drop(in_tx);
let mut buf = Vec::new();
socket.read_to_end(&mut buf).await.unwrap();
assert_eq!(buf, b"first second third");
assert_eq!(reads.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn frames_delivered_one_byte_at_a_time() {
let (mut sender, receiver) = transport_states();
let (
mock,
MockPeer {
inbound: in_tx,
outbound: _out_rx,
reads,
},
) = MockSocket::new();
let mut socket = NoiseSocket::new(mock, receiver);
let mut bytes = wire_frame(&mut sender, b"one byte");
bytes.extend(wire_frame(&mut sender, b" at a time"));
for byte in &bytes {
in_tx.send(vec![*byte]).unwrap();
}
drop(in_tx);
let mut buf = Vec::new();
socket.read_to_end(&mut buf).await.unwrap();
assert_eq!(buf, b"one byte at a time");
assert_eq!(reads.load(Ordering::SeqCst), bytes.len());
}
#[tokio::test]
async fn max_size_frames_larger_than_the_read_ahead_buffer() {
let (mut sender, receiver) = transport_states();
let (
mock,
MockPeer {
inbound: in_tx,
outbound: _out_rx,
reads: _reads,
},
) = MockSocket::new();
let mut socket = NoiseSocket::new(mock, receiver);
let mut chunk = Vec::new();
let mut expected = Vec::new();
for value in 1..=3u8 {
let message = vec![value; MAX_WRITE_BUFFER_LENGTH];
chunk.extend(wire_frame(&mut sender, &message));
expected.extend(message);
}
assert!(chunk.len() > READ_AHEAD_LENGTH);
in_tx.send(chunk).unwrap();
drop(in_tx);
let mut buf = Vec::new();
socket.read_to_end(&mut buf).await.unwrap();
assert_eq!(buf, expected);
}
#[tokio::test]
async fn handshake_end_and_transport_frames_in_one_read() {
let (initiator, responder) = handshake_states();
let mut initiator = NoiseState::from(initiator);
let (
mock,
MockPeer {
inbound: in_tx,
outbound: mut out_rx,
reads,
},
) = MockSocket::new();
let handshake = Handshake::new(mock, responder, Duration::from_secs(5));
let responder_task = tokio::spawn(handshake.perform_handshake());
in_tx.send(wire_frame(&mut initiator, &[])).unwrap();
let mut received = Vec::new();
loop {
received.extend(out_rx.recv().await.unwrap());
let frame_len = received
.get(..2)
.map(|len| usize::from(u16::from_be_bytes(len.try_into().unwrap())));
if let Some(message) = frame_len.and_then(|len| received.get(2..len.saturating_add(2))) {
let mut payload = vec![0u8; MAX_PAYLOAD_LENGTH];
initiator.read_message(message, &mut payload).unwrap();
break;
}
}
let mut chunk = wire_frame(&mut initiator, &[]);
let mut initiator = initiator.into_transport_mode().unwrap();
chunk.extend(wire_frame(&mut initiator, b"hello "));
chunk.extend(wire_frame(&mut initiator, b"world"));
in_tx.send(chunk).unwrap();
drop(in_tx);
let mut socket = responder_task.await.unwrap().unwrap();
let mut buf = Vec::new();
socket.read_to_end(&mut buf).await.unwrap();
assert_eq!(buf, b"hello world");
assert_eq!(reads.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn eof_inside_a_buffered_partial_frame_is_unexpected() {
let (mut sender, receiver) = transport_states();
let (
mock,
MockPeer {
inbound: in_tx,
outbound: _out_rx,
reads: _reads,
},
) = MockSocket::new();
let mut socket = NoiseSocket::new(mock, receiver);
let mut chunk = wire_frame(&mut sender, b"complete");
let truncated = wire_frame(&mut sender, b"truncated");
chunk.extend(truncated.iter().take(truncated.len().saturating_sub(3)));
in_tx.send(chunk).unwrap();
drop(in_tx);
let mut buf = [0u8; 8];
socket.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"complete");
let err = socket.read(&mut buf).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
assert!(socket.is_eof());
let err = socket.read(&mut buf).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
}
#[tokio::test]
async fn eof_inside_a_buffered_length_prefix_is_unexpected() {
let (mut sender, receiver) = transport_states();
let (
mock,
MockPeer {
inbound: in_tx,
outbound: _out_rx,
reads: _reads,
},
) = MockSocket::new();
let mut socket = NoiseSocket::new(mock, receiver);
let mut chunk = wire_frame(&mut sender, b"complete");
chunk.push(0);
in_tx.send(chunk).unwrap();
drop(in_tx);
let mut buf = [0u8; 8];
socket.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"complete");
let err = socket.read(&mut buf).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
}
#[tokio::test]
async fn empty_frame_flood_yields_to_the_scheduler() {
let (mut sender, receiver) = transport_states();
let (mock, MockPeer { inbound: in_tx, .. }) = MockSocket::new();
let mut socket = NoiseSocket::new(mock, receiver);
let mut chunk = vec![0u8; 200_000];
chunk.extend(wire_frame(&mut sender, b"after the flood"));
in_tx.send(chunk).unwrap();
drop(in_tx);
let mut buf = [0u8; 15];
{
let mut read = std::pin::pin!(socket.read(&mut buf));
assert!(futures::poll!(read.as_mut()).is_pending());
}
socket.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"after the flood");
}
#[tokio::test]
async fn buffered_frame_flood_yields_to_the_scheduler() {
let (mut sender, receiver) = transport_states();
let (mock, MockPeer { inbound: in_tx, .. }) = MockSocket::new();
let mut socket = NoiseSocket::new(mock, receiver);
let mut chunk = Vec::new();
for _ in 0..1000 {
chunk.extend(wire_frame(&mut sender, b"x"));
}
in_tx.send(chunk).unwrap();
drop(in_tx);
let frames_in_one_poll = std::future::poll_fn(|cx| {
let mut frames = 0usize;
loop {
let mut byte = [0u8; 1];
let mut read_buf = ReadBuf::new(&mut byte);
match Pin::new(&mut socket).poll_read(cx, &mut read_buf) {
Poll::Ready(result) => {
result.unwrap();
assert_eq!(read_buf.filled(), b"x");
frames = frames.saturating_add(1);
},
Poll::Pending => return Poll::Ready(frames),
}
}
})
.await;
assert!(frames_in_one_poll > 0);
assert!(
frames_in_one_poll < 100,
"{frames_in_one_poll} frames decrypted in one poll"
);
let mut rest = Vec::new();
socket.read_to_end(&mut rest).await.unwrap();
assert_eq!(frames_in_one_poll.saturating_add(rest.len()), 1000);
assert!(rest.iter().all(|b| *b == b'x'));
}
}