use std::cmp;
use std::io;
use std::pin::Pin;
use std::task::{Context, Poll};
use bytes::{Buf as _, BytesMut};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use zeroize::Zeroizing;
use crate::aead::AeadCipher;
use crate::method::CipherMethod;
use crate::nonce::NonceCounter;
pub const MAX_CHUNK_PAYLOAD: usize = 16 * 1024 - 1;
#[derive(Clone, Copy, Debug)]
enum ReadState {
PeerSalt,
LengthBlock,
Payload { len: usize },
Closed,
}
pub struct ShadowsocksAeadStream<S> {
inner: S,
method: CipherMethod,
write_subkey: Zeroizing<Vec<u8>>,
write_cipher: Option<AeadCipher>,
read_subkey: Option<Zeroizing<Vec<u8>>>,
read_cipher: Option<AeadCipher>,
password_ikm: Option<Zeroizing<Vec<u8>>>,
send_write_salt: bool,
write_nonce: NonceCounter,
read_nonce: NonceCounter,
read_plain: BytesMut,
read_buf: BytesMut,
read_state: ReadState,
write_buf: BytesMut,
}
impl<S: AsyncRead + AsyncWrite + Unpin> ShadowsocksAeadStream<S> {
pub fn new(
inner: S,
method: CipherMethod,
subkey: Vec<u8>,
) -> Result<Self, crate::error::ShadowsocksError> {
let nonce_size = method.nonce_size();
let write_cipher = AeadCipher::new(method, &subkey)?;
let read_cipher = AeadCipher::new(method, &subkey)?;
let subkey = Zeroizing::new(subkey);
Ok(Self {
inner,
method,
write_subkey: subkey.clone(),
write_cipher: Some(write_cipher),
read_subkey: Some(subkey),
read_cipher: Some(read_cipher),
password_ikm: None,
send_write_salt: false,
write_nonce: NonceCounter::starting_at(nonce_size, 0),
read_nonce: NonceCounter::starting_at(nonce_size, 0),
read_plain: BytesMut::new(),
read_buf: BytesMut::new(),
read_state: ReadState::LengthBlock,
write_buf: BytesMut::new(),
})
}
pub fn new_client(
inner: S,
method: CipherMethod,
write_subkey: Vec<u8>,
password: &str,
) -> Result<Self, crate::error::ShadowsocksError> {
let nonce_size = method.nonce_size();
let write_cipher = AeadCipher::new(method, &write_subkey)?;
Ok(Self {
inner,
method,
write_subkey: Zeroizing::new(write_subkey),
write_cipher: Some(write_cipher),
read_subkey: None,
read_cipher: None,
password_ikm: Some(Zeroizing::new(CipherMethod::password_key_material(
password.as_bytes(),
))),
send_write_salt: false,
write_nonce: NonceCounter::starting_at(nonce_size, 2),
read_nonce: NonceCounter::starting_at(nonce_size, 0),
read_plain: BytesMut::new(),
read_buf: BytesMut::new(),
read_state: ReadState::PeerSalt,
write_buf: BytesMut::new(),
})
}
pub fn new_server(
inner: S,
method: CipherMethod,
read_subkey: Vec<u8>,
send_write_salt: bool,
password: &str,
) -> Result<Self, crate::error::ShadowsocksError> {
if !send_write_salt {
return Err(crate::error::ShadowsocksError::Other(
"server streams must send a write salt".to_string(),
));
}
let nonce_size = method.nonce_size();
let read_cipher = AeadCipher::new(method, &read_subkey)?;
Ok(Self {
inner,
method,
write_subkey: Zeroizing::new(Vec::new()), write_cipher: None,
read_subkey: Some(Zeroizing::new(read_subkey)),
read_cipher: Some(read_cipher),
password_ikm: Some(Zeroizing::new(CipherMethod::password_key_material(
password.as_bytes(),
))),
send_write_salt,
read_nonce: NonceCounter::starting_at(nonce_size, 2),
write_nonce: NonceCounter::starting_at(nonce_size, 0),
read_plain: BytesMut::new(),
read_buf: BytesMut::new(),
read_state: ReadState::LengthBlock,
write_buf: BytesMut::new(),
})
}
pub fn new_with_nonces(
inner: S,
method: CipherMethod,
subkey: Vec<u8>,
write_start: u64,
read_start: u64,
) -> Result<Self, crate::error::ShadowsocksError> {
let nonce_size = method.nonce_size();
let write_cipher = AeadCipher::new(method, &subkey)?;
let read_cipher = AeadCipher::new(method, &subkey)?;
let subkey = Zeroizing::new(subkey);
Ok(Self {
inner,
method,
write_subkey: subkey.clone(),
write_cipher: Some(write_cipher),
read_subkey: Some(subkey),
read_cipher: Some(read_cipher),
password_ikm: None,
send_write_salt: false,
write_nonce: NonceCounter::starting_at(nonce_size, write_start),
read_nonce: NonceCounter::starting_at(nonce_size, read_start),
read_plain: BytesMut::new(),
read_buf: BytesMut::new(),
read_state: ReadState::LengthBlock,
write_buf: BytesMut::new(),
})
}
pub fn into_inner(self) -> S {
self.inner
}
pub(crate) fn prepend_read_plaintext(&mut self, plaintext: &[u8]) {
if !plaintext.is_empty() {
self.read_plain.extend_from_slice(plaintext);
}
}
}
fn read_until<S: AsyncRead + Unpin>(
inner: &mut S,
cx: &mut Context<'_>,
buf: &mut BytesMut,
target: usize,
) -> Poll<io::Result<bool>> {
while buf.len() < target {
let start = buf.len();
buf.resize(target, 0);
let mut rbuf = ReadBuf::new(&mut buf[start..target]);
match Pin::new(&mut *inner).poll_read(cx, &mut rbuf) {
Poll::Ready(Ok(())) => {
let n = rbuf.filled().len();
if n == 0 {
buf.truncate(start);
return Poll::Ready(Ok(false));
}
buf.truncate(start + n);
}
Poll::Ready(Err(e)) => {
buf.truncate(start);
return Poll::Ready(Err(e));
}
Poll::Pending => {
buf.truncate(start);
return Poll::Pending;
}
}
}
Poll::Ready(Ok(true))
}
impl<S: AsyncRead + AsyncWrite + Unpin> AsyncRead for ShadowsocksAeadStream<S> {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let this = self.get_mut();
if !this.read_plain.is_empty() {
let n = cmp::min(this.read_plain.len(), buf.remaining());
buf.put_slice(&this.read_plain.split_to(n));
return Poll::Ready(Ok(()));
}
loop {
let state = this.read_state;
match state {
ReadState::Closed => return Poll::Ready(Ok(())),
ReadState::PeerSalt => {
let salt_size = this.method.salt_size();
match read_until(&mut this.inner, cx, &mut this.read_buf, salt_size) {
Poll::Ready(Ok(true)) => {}
Poll::Ready(Ok(false)) => {
this.read_buf.clear();
return Poll::Ready(Ok(()));
}
Poll::Ready(Err(e)) => {
this.read_buf.clear();
return Poll::Ready(Err(e));
}
Poll::Pending => return Poll::Pending,
}
let salt = this.read_buf.split_to(salt_size);
let password_ikm = this
.password_ikm
.as_deref()
.ok_or_else(|| io::Error::other("no password for subkey derivation"))?;
let read_subkey = this
.method
.derive_key_from_ikm(password_ikm, &salt)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
let read_cipher = AeadCipher::new(this.method, &read_subkey)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
this.read_subkey = Some(Zeroizing::new(read_subkey));
this.read_cipher = Some(read_cipher);
this.read_buf.clear();
this.read_state = ReadState::LengthBlock;
}
ReadState::LengthBlock => {
let len_block_size = 2 + this.method.tag_size();
match read_until(&mut this.inner, cx, &mut this.read_buf, len_block_size) {
Poll::Ready(Ok(true)) => {}
Poll::Ready(Ok(false)) => {
this.read_buf.clear();
return Poll::Ready(Ok(()));
}
Poll::Ready(Err(e)) => {
this.read_buf.clear();
return Poll::Ready(Err(e));
}
Poll::Pending => return Poll::Pending,
}
let cipher = this
.read_cipher
.as_ref()
.ok_or_else(|| io::Error::other("read cipher not yet derived"))?;
let mut nonce = [0u8; 12];
this.read_nonce
.current(&mut nonce)
.map_err(io::Error::other)?;
let len_plaintext = cipher
.decrypt(&nonce, &this.read_buf[..len_block_size])
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
this.read_nonce.advance().map_err(io::Error::other)?;
if len_plaintext.len() != 2 {
this.read_buf.clear();
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidData,
"invalid length block plaintext",
)));
}
let payload_len =
u16::from_be_bytes([len_plaintext[0], len_plaintext[1]]) as usize;
this.read_buf.clear();
if payload_len > MAX_CHUNK_PAYLOAD {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"payload length {} exceeds maximum {}",
payload_len, MAX_CHUNK_PAYLOAD
),
)));
}
if payload_len == 0 {
this.read_state = ReadState::Closed;
return Poll::Ready(Ok(()));
}
this.read_state = ReadState::Payload { len: payload_len };
}
ReadState::Payload { len } => {
let wire_len = len + this.method.tag_size();
match read_until(&mut this.inner, cx, &mut this.read_buf, wire_len) {
Poll::Ready(Ok(true)) => {}
Poll::Ready(Ok(false)) | Poll::Ready(Err(_)) => {
this.read_buf.clear();
this.read_state = ReadState::LengthBlock;
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"unexpected EOF in payload",
)));
}
Poll::Pending => return Poll::Pending,
}
let cipher = this
.read_cipher
.as_ref()
.ok_or_else(|| io::Error::other("read cipher not yet derived"))?;
let mut nonce = [0u8; 12];
this.read_nonce
.current(&mut nonce)
.map_err(io::Error::other)?;
let plaintext = cipher
.decrypt(&nonce, &this.read_buf)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
this.read_nonce.advance().map_err(io::Error::other)?;
this.read_buf.clear();
this.read_state = ReadState::LengthBlock;
if plaintext.len() <= buf.remaining() {
buf.put_slice(&plaintext);
} else {
this.read_plain.extend_from_slice(&plaintext);
let n = cmp::min(this.read_plain.len(), buf.remaining());
buf.put_slice(&this.read_plain.split_to(n));
}
return Poll::Ready(Ok(()));
}
}
}
}
}
impl<S: AsyncRead + AsyncWrite + Unpin> AsyncWrite for ShadowsocksAeadStream<S> {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let this = self.get_mut();
if this.send_write_salt {
use rand::RngCore;
let mut salt_buf = [0u8; 32];
rand::thread_rng().fill_bytes(&mut salt_buf[..this.method.salt_size()]);
let salt = &salt_buf[..this.method.salt_size()];
let password_ikm = this
.password_ikm
.as_deref()
.ok_or_else(|| io::Error::other("no password for subkey derivation"))?;
let write_subkey = this
.method
.derive_key_from_ikm(password_ikm, salt)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
let write_cipher = AeadCipher::new(this.method, &write_subkey)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
this.write_subkey = Zeroizing::new(write_subkey);
this.write_cipher = Some(write_cipher);
this.write_buf.extend_from_slice(salt);
this.send_write_salt = false;
}
while !this.write_buf.is_empty() {
match Pin::new(&mut this.inner).poll_write(cx, &this.write_buf) {
Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::WriteZero,
"zero-byte write",
)));
}
Poll::Ready(Ok(n)) => this.write_buf.advance(n),
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => return Poll::Pending,
}
}
let chunk_size = cmp::min(buf.len(), MAX_CHUNK_PAYLOAD);
if chunk_size == 0 {
return Poll::Ready(Ok(0));
}
let len_bytes = (chunk_size as u16).to_be_bytes();
let cipher = this
.write_cipher
.as_ref()
.ok_or_else(|| io::Error::other("write cipher not yet derived"))?;
let mut len_nonce = [0u8; 12];
this.write_nonce
.current(&mut len_nonce)
.map_err(io::Error::other)?;
let len_ct = cipher
.encrypt(&len_nonce, &len_bytes)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
this.write_nonce.advance().map_err(io::Error::other)?;
let mut payload_nonce = [0u8; 12];
this.write_nonce
.current(&mut payload_nonce)
.map_err(io::Error::other)?;
let payload_ct = cipher
.encrypt(&payload_nonce, &buf[..chunk_size])
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
this.write_nonce.advance().map_err(io::Error::other)?;
this.write_buf.extend_from_slice(&len_ct);
this.write_buf.extend_from_slice(&payload_ct);
while !this.write_buf.is_empty() {
match Pin::new(&mut this.inner).poll_write(cx, &this.write_buf) {
Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::WriteZero,
"zero-byte write",
)));
}
Poll::Ready(Ok(n)) => this.write_buf.advance(n),
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => break,
}
}
Poll::Ready(Ok(chunk_size))
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
while !this.write_buf.is_empty() {
match Pin::new(&mut this.inner).poll_write(cx, &this.write_buf) {
Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::WriteZero,
"zero-byte write",
)));
}
Poll::Ready(Ok(n)) => this.write_buf.advance(n),
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => return Poll::Pending,
}
}
Pin::new(&mut this.inner).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
while !this.write_buf.is_empty() {
match Pin::new(&mut this.inner).poll_write(cx, &this.write_buf) {
Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::WriteZero,
"zero-byte write",
)));
}
Poll::Ready(Ok(n)) => this.write_buf.advance(n),
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => return Poll::Pending,
}
}
Pin::new(&mut this.inner).poll_shutdown(cx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use crate::aead::encrypt_chunk_standard;
#[test]
fn invalid_subkey_returns_error_instead_of_panicking() {
let (stream, _) = tokio::io::duplex(64);
let result = ShadowsocksAeadStream::new(stream, CipherMethod::Aes256Gcm, vec![0; 16]);
assert!(result.is_err());
}
#[tokio::test]
async fn roundtrip_small_data() {
let (client, server) = tokio::io::duplex(4096);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x42u8; 32];
let mut client_stream = ShadowsocksAeadStream::new(client, method, subkey.clone()).unwrap();
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey).unwrap();
client_stream.write_all(b"hello").await.unwrap();
client_stream.flush().await.unwrap();
let mut buf = vec![0u8; 64];
let n = server_stream.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"hello");
}
#[tokio::test]
async fn roundtrip_large_data() {
let (client, server) = tokio::io::duplex(1 << 16);
let method = CipherMethod::ChaCha20IetfPoly1305;
let subkey = vec![0xABu8; 32];
let payload = vec![0xCDu8; 100_000];
let expected = payload.clone();
let write_subkey = subkey.clone();
let write_handle = tokio::spawn(async move {
let mut client_stream =
ShadowsocksAeadStream::new(client, method, write_subkey).unwrap();
client_stream.write_all(&payload).await.unwrap();
client_stream.flush().await.unwrap();
});
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey).unwrap();
let mut received = Vec::new();
server_stream.read_to_end(&mut received).await.unwrap();
write_handle.await.unwrap();
assert_eq!(received, expected);
}
#[tokio::test]
async fn bidirectional_communication() {
let (c1, s1) = tokio::io::duplex(4096);
let (c2, s2) = tokio::io::duplex(4096);
let method = CipherMethod::Aes128Gcm;
let subkey = vec![0x11u8; 16];
let mut client_a = ShadowsocksAeadStream::new(c1, method, subkey.clone()).unwrap();
let mut server_a = ShadowsocksAeadStream::new(s1, method, subkey.clone()).unwrap();
let mut client_b = ShadowsocksAeadStream::new(c2, method, subkey.clone()).unwrap();
let mut server_b = ShadowsocksAeadStream::new(s2, method, subkey).unwrap();
client_a.write_all(b"ping").await.unwrap();
client_a.flush().await.unwrap();
let mut buf = vec![0u8; 64];
let n = server_a.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"ping");
client_b.write_all(b"pong").await.unwrap();
client_b.flush().await.unwrap();
let n = server_b.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"pong");
}
#[tokio::test]
async fn empty_read_on_eof() {
let (client, server) = tokio::io::duplex(256);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x55u8; 32];
let client_stream = ShadowsocksAeadStream::new(client, method, subkey.clone()).unwrap();
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey).unwrap();
drop(client_stream);
let mut buf = vec![0u8; 64];
let result = server_stream.read(&mut buf).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), 0); }
#[tokio::test]
async fn write_buffer_flushed_on_flush() {
let (client, server) = tokio::io::duplex(4096);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x99u8; 32];
let mut client_stream = ShadowsocksAeadStream::new(client, method, subkey.clone()).unwrap();
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey).unwrap();
client_stream.write_all(b"data").await.unwrap();
client_stream.flush().await.unwrap();
let mut buf = vec![0u8; 64];
let n = server_stream.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"data");
}
#[tokio::test]
async fn multiple_chunks() {
let (client, server) = tokio::io::duplex(8192);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x77u8; 32];
let mut client_stream = ShadowsocksAeadStream::new(client, method, subkey.clone()).unwrap();
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey).unwrap();
for i in 0..10 {
let msg = format!("msg-{i}");
client_stream.write_all(msg.as_bytes()).await.unwrap();
}
client_stream.flush().await.unwrap();
drop(client_stream);
let mut received = Vec::new();
server_stream.read_to_end(&mut received).await.unwrap();
let mut expected = Vec::new();
for i in 0..10 {
expected.extend_from_slice(format!("msg-{i}").as_bytes());
}
assert_eq!(received, expected);
}
#[tokio::test]
async fn into_inner_returns_original_stream() {
let (client, _server) = tokio::io::duplex(256);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x01u8; 32];
let stream = ShadowsocksAeadStream::new(client, method, subkey).unwrap();
let _ = stream.into_inner();
}
#[tokio::test]
async fn zero_length_payload_signals_eof() {
let (client, server) = tokio::io::duplex(256);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x33u8; 32];
let mut server_stream =
ShadowsocksAeadStream::new_with_nonces(server, method, subkey.clone(), 0, 0).unwrap();
let mut nonce = vec![0u8; method.nonce_size()];
nonce[0] = 0;
let wire = encrypt_chunk_standard(method, &subkey, &nonce, b"").unwrap();
let mut raw_stream = client;
raw_stream.write_all(&wire).await.unwrap();
let mut buf = vec![0u8; 64];
let result = server_stream.read(&mut buf).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), 0);
raw_stream.write_all(b"ignored").await.unwrap();
let result = server_stream.read(&mut buf).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), 0);
}
#[tokio::test]
async fn tampered_length_block_fails() {
let (client, server) = tokio::io::duplex(256);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x42u8; 32];
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey.clone()).unwrap();
let mut raw_stream = client;
let nonce1 = {
let mut n = vec![0u8; method.nonce_size()];
n[0] = 2; n
};
let len_bytes = (5u16).to_be_bytes();
let len_ct = crate::aead::aead_encrypt_raw(method, &subkey, &nonce1, &len_bytes).unwrap();
let mut tampered_len_ct = len_ct;
tampered_len_ct[0] ^= 0xFF;
raw_stream.write_all(&tampered_len_ct).await.unwrap();
let nonce2 = {
let mut n = vec![0u8; method.nonce_size()];
n[0] = 3; n
};
let payload_ct = crate::aead::aead_encrypt_raw(method, &subkey, &nonce2, b"hello").unwrap();
raw_stream.write_all(&payload_ct).await.unwrap();
drop(raw_stream);
let mut buf = vec![0u8; 64];
let result = tokio::time::timeout(
std::time::Duration::from_secs(2),
server_stream.read(&mut buf),
)
.await;
let result = result.expect("tampered length block must fail promptly");
let error = result.expect_err("tampered length block must not produce plaintext or EOF");
assert_eq!(error.kind(), std::io::ErrorKind::InvalidData);
}
#[test]
fn server_without_write_salt_is_rejected() {
let (_, server) = tokio::io::duplex(64);
let result = ShadowsocksAeadStream::new_server(
server,
CipherMethod::Aes256Gcm,
vec![0x42; 32],
false,
"password",
);
assert!(result.is_err());
}
#[tokio::test]
async fn tampered_payload_block_fails() {
let (client, server) = tokio::io::duplex(1024);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x42u8; 32];
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey.clone()).unwrap();
let mut raw_stream = client;
let nonce1 = {
let mut n = vec![0u8; method.nonce_size()];
n[0] = 2; n
};
let len_bytes = (5u16).to_be_bytes();
let len_ct = crate::aead::aead_encrypt_raw(method, &subkey, &nonce1, &len_bytes).unwrap();
raw_stream.write_all(&len_ct).await.unwrap();
let nonce2 = {
let mut n = vec![0u8; method.nonce_size()];
n[0] = 3; n
};
let payload_ct = crate::aead::aead_encrypt_raw(method, &subkey, &nonce2, b"hello").unwrap();
let mut tampered_payload_ct = payload_ct;
tampered_payload_ct[0] ^= 0xFF;
raw_stream.write_all(&tampered_payload_ct).await.unwrap();
drop(raw_stream);
let mut buf = vec![0u8; 64];
let result = tokio::time::timeout(
std::time::Duration::from_secs(2),
server_stream.read(&mut buf),
)
.await;
let result = result.expect("tampered payload block must fail promptly");
let error = result.expect_err("tampered payload block must not produce plaintext or EOF");
assert_eq!(error.kind(), std::io::ErrorKind::InvalidData);
}
#[tokio::test]
async fn tampered_payload_fails() {
use crate::aead::aead_encrypt_raw;
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x42u8; 32];
let plaintext = b"hello world";
let nonce = vec![0u8; method.nonce_size()];
let ciphertext = aead_encrypt_raw(method, &subkey, &nonce, plaintext).unwrap();
let mut tampered = ciphertext.clone();
tampered[0] ^= 0x01;
let result = crate::aead::aead_decrypt_raw(method, &subkey, &nonce, &tampered);
assert!(
result.is_err(),
"decryption of tampered ciphertext should fail"
);
}
#[tokio::test]
async fn wrong_key_fails() {
use crate::aead::aead_encrypt_raw;
let method = CipherMethod::Aes256Gcm;
let correct_key = vec![0x42u8; 32];
let wrong_key = vec![0x99u8; 32];
let plaintext = b"secret data";
let nonce = vec![0u8; method.nonce_size()];
let ciphertext = aead_encrypt_raw(method, &correct_key, &nonce, plaintext).unwrap();
let result = crate::aead::aead_decrypt_raw(method, &wrong_key, &nonce, &ciphertext);
assert!(result.is_err(), "decryption with wrong key should fail");
let result = crate::aead::aead_decrypt_raw(method, &correct_key, &nonce, &ciphertext);
assert!(result.is_ok(), "decryption with correct key should succeed");
assert_eq!(result.unwrap(), plaintext);
}
#[tokio::test]
async fn standard_chunk_format_roundtrip() {
let (client, server) = tokio::io::duplex(4096);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0x55u8; 32];
let mut client_stream = ShadowsocksAeadStream::new(client, method, subkey.clone()).unwrap();
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey).unwrap();
let data = b"standard SIP003 framing test";
client_stream.write_all(data).await.unwrap();
client_stream.flush().await.unwrap();
let mut buf = vec![0u8; 64];
let n = server_stream.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], data.as_slice());
}
#[tokio::test]
async fn empty_plaintext_roundtrip() {
let (client, server) = tokio::io::duplex(256);
let method = CipherMethod::Aes128Gcm;
let subkey = vec![0x88u8; 16];
let mut client_stream = ShadowsocksAeadStream::new(client, method, subkey.clone()).unwrap();
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey).unwrap();
client_stream.write_all(b"").await.unwrap();
client_stream.flush().await.unwrap();
drop(client_stream);
let mut buf = vec![0u8; 64];
let result = server_stream.read(&mut buf).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), 0); }
#[tokio::test]
async fn large_chunk_split_across_reads() {
let (client, server) = tokio::io::duplex(1 << 18);
let method = CipherMethod::Aes256Gcm;
let subkey = vec![0xBBu8; 32];
let payload = vec![0x44u8; 50_000];
let expected = payload.clone();
let write_subkey = subkey.clone();
let write_handle = tokio::spawn(async move {
let mut client_stream =
ShadowsocksAeadStream::new(client, method, write_subkey).unwrap();
client_stream.write_all(&payload).await.unwrap();
client_stream.flush().await.unwrap();
drop(client_stream);
});
let mut server_stream = ShadowsocksAeadStream::new(server, method, subkey).unwrap();
let mut received = Vec::new();
let mut tmp = [0u8; 1024];
loop {
let n = server_stream.read(&mut tmp).await.unwrap();
if n == 0 {
break;
}
received.extend_from_slice(&tmp[..n]);
}
write_handle.await.unwrap();
assert_eq!(received, expected);
}
}