use std::time::Duration;
use tari_test_utils::unpack_enum;
use tokio::{task, time};
use crate::{
framing,
memsocket::MemorySocket,
protocol::rpc::{
Handshake,
RPC_MAX_FRAME_SIZE,
error::HandshakeRejectReason,
handshake::{MAX_HANDSHAKE_FRAME_SIZE, RpcHandshakeError, SUPPORTED_RPC_VERSIONS},
},
};
fn oversized_handshake_frame() -> bytes::Bytes {
let mut frame = Vec::new();
prost::encoding::encode_key(15, prost::encoding::WireType::LengthDelimited, &mut frame);
prost::encoding::encode_varint((MAX_HANDSHAKE_FRAME_SIZE - 2) as u64, &mut frame);
frame.resize(MAX_HANDSHAKE_FRAME_SIZE + 1, 0);
frame.into()
}
#[tokio::test]
async fn the_server_rejects_an_oversized_handshake_frame() {
use futures::SinkExt;
let (client, server) = MemorySocket::new_pair();
let mut client_framed = framing::canonical(client, 4096);
client_framed.send(oversized_handshake_frame()).await.unwrap();
let mut server_framed = framing::canonical(server, 4096);
let err = Handshake::new(&mut server_framed)
.perform_server_handshake()
.await
.unwrap_err();
unpack_enum!(RpcHandshakeError::FrameTooLarge { max } = err);
assert_eq!(max, MAX_HANDSHAKE_FRAME_SIZE);
assert_eq!(server_framed.codec().max_frame_length(), 4096);
}
#[tokio::test]
async fn a_handshake_frame_declaring_more_than_the_limit_is_rejected_by_the_codec() {
use tokio::io::AsyncWriteExt;
let declared = u32::try_from(RPC_MAX_FRAME_SIZE).unwrap();
for server_side in [true, false] {
let (mut peer, local) = MemorySocket::new_pair();
peer.write_all(&declared.to_be_bytes()).await.unwrap();
let mut framed = framing::canonical(local, RPC_MAX_FRAME_SIZE);
let result = time::timeout(Duration::from_secs(5), async {
let mut handshake = Handshake::new(&mut framed);
if server_side {
handshake.perform_server_handshake().await.map(|_| ())
} else {
handshake.perform_client_handshake().await
}
})
.await
.expect("the codec did not reject the frame from its declared length");
unpack_enum!(RpcHandshakeError::FrameTooLarge { max } = result.unwrap_err());
assert_eq!(max, MAX_HANDSHAKE_FRAME_SIZE);
assert_eq!(framed.codec().max_frame_length(), RPC_MAX_FRAME_SIZE);
drop(peer);
}
}
#[tokio::test]
async fn the_client_rejects_an_oversized_handshake_reply() {
use futures::SinkExt;
let (client, server) = MemorySocket::new_pair();
let mut server_framed = framing::canonical(server, 4096);
server_framed.send(oversized_handshake_frame()).await.unwrap();
let mut client_framed = framing::canonical(client, 4096);
let err = Handshake::new(&mut client_framed)
.perform_client_handshake()
.await
.unwrap_err();
unpack_enum!(RpcHandshakeError::FrameTooLarge { .. } = err);
assert!(crate::protocol::rpc::RpcError::from(err).is_caused_by_server());
}
#[tokio::test]
async fn it_performs_the_handshake() {
let (client, server) = MemorySocket::new_pair();
let handshake_result = task::spawn(async move {
let mut server_framed = framing::canonical(server, 1024);
let mut handshake_server = Handshake::new(&mut server_framed);
handshake_server.perform_server_handshake().await
});
let mut client_framed = framing::canonical(client, 1024);
let mut handshake_client = Handshake::new(&mut client_framed);
handshake_client.perform_client_handshake().await.unwrap();
let v = handshake_result.await.unwrap().unwrap();
assert!(SUPPORTED_RPC_VERSIONS.contains(&v));
}
#[tokio::test]
async fn it_rejects_the_handshake() {
let (client, server) = MemorySocket::new_pair();
let mut client_framed = framing::canonical(client, 1024);
let mut handshake_client = Handshake::new(&mut client_framed);
let mut server_framed = framing::canonical(server, 1024);
let mut handshake_server = Handshake::new(&mut server_framed);
handshake_server
.reject_with_reason(HandshakeRejectReason::NoServerSessionsAvailable("some reason"))
.await
.unwrap();
let err = handshake_client.perform_client_handshake().await.unwrap_err();
unpack_enum!(RpcHandshakeError::Rejected(reason) = err);
unpack_enum!(HandshakeRejectReason::NoServerSessionsAvailable("session limit reached") = reason);
}