#![forbid(unsafe_code)]
#![warn(clippy::all)]
mod helper_buf;
use chacha20poly1305::ChaCha20Poly1305;
use chacha20poly1305::aead::Buffer;
use chacha20poly1305::aead::stream::{DecryptorBE32, EncryptorBE32};
use helper_buf::HelperBuf;
use pin_project::pin_project;
use std::io::ErrorKind;
use std::pin::Pin;
use std::task::{Context, Poll, ready};
use tokio::io::{AsyncBufRead, AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
const TAG_SIZE: usize = 16;
#[pin_project]
pub struct EncryptedStream<T> {
#[pin]
inner: T,
decryptor: DecryptorBE32<ChaCha20Poly1305>,
encryptor: EncryptorBE32<ChaCha20Poly1305>,
received: HelperBuf,
decrypted: HelperBuf,
to_send: HelperBuf,
flushing: bool,
}
impl<T> EncryptedStream<T> {
pub fn new(io_stream: T, key: &[u8; 32], nonce: &[u8; 7]) -> Self {
let mut to_send = HelperBuf::with_capacity(u16::MAX as usize + 2);
to_send.extend_from_slice(&[0, 0]).expect("unreachable");
Self {
inner: io_stream,
decryptor: DecryptorBE32::new(key.into(), nonce.into()),
encryptor: EncryptorBE32::new(key.into(), nonce.into()),
received: HelperBuf::with_capacity(u16::MAX as usize + 2),
decrypted: HelperBuf::with_capacity(u16::MAX as usize + 2),
to_send,
flushing: false,
}
}
}
impl<T: AsyncRead + AsyncWrite + Unpin> EncryptedStream<T> {
pub async fn encrypt_connection(
mut io_stream: T,
shared_key: &[u8; 32],
) -> std::io::Result<Self> {
let my_seed: [u8; 7] = rand::random();
io_stream.write_all(&my_seed).await?;
io_stream.flush().await?;
let mut peer_seed = [0; 7];
io_stream.read_exact(&mut peer_seed).await?;
peer_seed
.iter_mut()
.zip(my_seed.iter())
.for_each(|(x1, x2)| *x1 ^= *x2);
Ok(Self::new(io_stream, shared_key, &peer_seed))
}
}
impl<T: AsyncRead> AsyncRead for EncryptedStream<T> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
if self.decrypted.is_empty() {
ready!(self.as_mut().inner_read(cx))?;
}
let me = self.project();
let num_bytes = std::cmp::min(me.decrypted.len(), buf.remaining());
buf.put_slice(&me.decrypted[0..num_bytes]);
me.decrypted.consume(num_bytes);
Poll::Ready(Ok(()))
}
}
impl<T: AsyncRead> AsyncBufRead for EncryptedStream<T> {
fn consume(self: std::pin::Pin<&mut EncryptedStream<T>>, amt: usize) {
self.project().decrypted.consume(amt);
}
fn poll_fill_buf(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<std::io::Result<&[u8]>> {
if self.decrypted.is_empty() {
ready!(self.as_mut().inner_read(cx))?;
}
Poll::Ready(Ok(self.project().decrypted))
}
}
impl<T: AsyncWrite> AsyncWrite for EncryptedStream<T> {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
if self.flushing {
ready!(self.as_mut().flush_write_buf(cx))?;
}
let me = self.as_mut().project();
let bytes_taken = std::cmp::min(buf.len(), me.to_send.spare_capacity().len() - TAG_SIZE);
me.to_send
.extend_from_slice(&buf[0..bytes_taken])
.expect("unreachable");
if me.to_send.spare_capacity().len() - TAG_SIZE == 0 {
let _ = self.flush_write_buf(cx)?;
}
Poll::Ready(Ok(bytes_taken))
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
ready!(self.as_mut().flush_write_buf(cx))?;
self.project().inner.poll_flush(cx)
}
fn poll_shutdown(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
ready!(self.as_mut().poll_flush(cx))?;
self.project().inner.poll_shutdown(cx)
}
}
impl<T: AsyncRead> EncryptedStream<T> {
fn inner_read(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
let mut me = self.project();
debug_assert!(me.decrypted.is_empty());
me.received.left_align();
fn peek_cipher_chunk(data: &[u8]) -> Option<&[u8]> {
let len: [u8; 2] = data.get(0..2)?.try_into().expect("unreachable");
let len = u16::from_be_bytes(len) as usize;
data.get(2..2 + len)
}
while peek_cipher_chunk(me.received).is_none() {
let mut read_buf = ReadBuf::new(me.received.spare_capacity());
ready!(me.inner.as_mut().poll_read(cx, &mut read_buf))?;
let bytes_read = read_buf.filled().len();
if bytes_read == 0 {
if me.received.is_empty() {
return Poll::Ready(Ok(()));
} else {
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"Unexpected EOF within encrypted chunk.",
)));
}
}
me.received.increase_len(bytes_read);
}
while let Some(cipher_chunk) = peek_cipher_chunk(me.received) {
let mut decryption_space = me.decrypted.split_off_aead_buf(me.decrypted.len());
decryption_space
.extend_from_slice(cipher_chunk)
.expect("Unreachable");
me.received.consume(cipher_chunk.len() + 2);
me.decryptor
.decrypt_next_in_place(&[], &mut decryption_space)
.map_err(|_| std::io::Error::new(ErrorKind::InvalidData, "Decryption error"))?;
}
Poll::Ready(Ok(()))
}
}
impl<T: AsyncWrite> EncryptedStream<T> {
fn flush_write_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
let mut me = self.project();
if !*me.flushing {
*me.flushing = true;
let mut msg = me.to_send.split_off_aead_buf(2);
me.encryptor
.encrypt_next_in_place(&[], &mut msg)
.map_err(|_| std::io::Error::new(ErrorKind::InvalidData, "Encryption error"))?;
let len = u16::try_from(msg.len())
.expect("unreachable: Length of message buffer should always fit in u16")
.to_be_bytes();
me.to_send[0..2].copy_from_slice(&len);
}
while !me.to_send.is_empty() {
let bytes_written = ready!(me.inner.as_mut().poll_write(cx, me.to_send))?;
me.to_send.consume(bytes_written);
}
*me.flushing = false;
me.to_send
.extend_from_slice(&[0, 0])
.expect("unreachable: to_send must have space for the header.");
Poll::Ready(Ok(()))
}
}