use super::handshake::{client_handshake, server_handshake};
use crate::client::{Config as ClientConfig, Connection as Client, ReconStrat};
use crate::error::TaskError;
use crate::handler::handler::{
Handler, PacketStream, bg_stream, bg_stream_reconnect,
};
use crate::handler::{Configurator, TaskHandle};
use crate::packet::builder::{PacketReceiver, PacketReceiverError};
use crate::packet::{EncryptedBytes, Packet};
use crate::server::{Config as ServerConfig, Connection as Server};
use crate::util::{ByteStream, TimeoutReader};
use std::io;
use tokio::io::AsyncWriteExt;
use tokio::sync::oneshot;
use tokio::time::Duration;
use crypto::cipher::Key;
use crypto::signature as sign;
pub fn client<S, P>(
stream: S,
cfg: ClientConfig,
recon_strat: Option<ReconStrat<S>>,
sign: sign::PublicKey,
) -> Client<P>
where
S: ByteStream,
P: Packet<EncryptedBytes> + Send + 'static,
P::Header: Send,
{
let (sender, receiver, mut bg_handler) = Handler::new(false);
let (cfg_tx, mut cfg_rx) = Configurator::new(cfg);
let (tx_close, mut rx_close) = oneshot::channel();
let task = tokio::spawn(async move {
bg_stream_reconnect(
stream,
&mut bg_handler,
&mut cfg_rx,
&mut rx_close,
|stream: &mut EncryptedPacketStream<_, _>, cfg| {
stream.stream.set_timeout(cfg.timeout);
stream.builder.set_body_limit(cfg.body_limit);
},
recon_strat,
|stream, cfg| {
EncryptedPacketStream::client(stream, sign.clone(), cfg)
},
)
.await
});
let task = TaskHandle {
close: tx_close,
task,
};
Client::new_raw(sender, receiver, cfg_tx, task)
}
pub fn server<S, P>(
stream: S,
cfg: ServerConfig,
sign: sign::Keypair,
) -> Server<P>
where
S: ByteStream,
P: Packet<EncryptedBytes> + Send + 'static,
P::Header: Send,
{
let (sender, receiver, mut bg_handler) = Handler::new(true);
let (cfg_tx, mut cfg_rx) = Configurator::new(cfg);
let (tx_close, mut rx_close) = oneshot::channel();
let task = tokio::spawn(async move {
let stream =
EncryptedPacketStream::server(stream, sign, cfg_rx.newest())
.await?;
let r = bg_stream(
stream,
&mut bg_handler,
&mut cfg_rx,
&mut rx_close,
|stream, cfg| {
stream.stream.set_timeout(cfg.timeout);
stream.builder.set_body_limit(cfg.body_limit);
},
)
.await;
if let Err(e) = &r {
tracing::error!("bg_stream closed with error {:?}", e);
}
r
});
let task = TaskHandle {
close: tx_close,
task,
};
Server::new_raw(sender, receiver, cfg_tx, task)
}
struct EncryptedPacketStream<S, P>
where
S: ByteStream,
P: Packet<EncryptedBytes>,
{
stream: TimeoutReader<S>,
send_key: Key,
recv_key: Key,
builder: PacketReceiver<P, EncryptedBytes>,
}
impl<S, P> EncryptedPacketStream<S, P>
where
S: ByteStream,
P: Packet<EncryptedBytes>,
{
async fn client(
stream: S,
sign: sign::PublicKey,
cfg: ClientConfig,
) -> Result<Self, TaskError> {
let mut stream = TimeoutReader::new(stream, cfg.timeout);
let handshake = client_handshake(&sign, &mut stream).await?;
Ok(Self {
stream,
send_key: handshake.send_key,
recv_key: handshake.recv_key,
builder: PacketReceiver::new(cfg.body_limit),
})
}
async fn server(
stream: S,
sign: sign::Keypair,
cfg: ServerConfig,
) -> Result<Self, TaskError> {
let mut stream = TimeoutReader::new(stream, cfg.timeout);
let handshake = server_handshake(&sign, &mut stream).await?;
Ok(Self {
stream,
send_key: handshake.send_key,
recv_key: handshake.recv_key,
builder: PacketReceiver::new(cfg.body_limit),
})
}
}
impl<S, P> PacketStream<P, EncryptedBytes> for EncryptedPacketStream<S, P>
where
S: ByteStream,
P: Packet<EncryptedBytes>,
{
fn timeout(&self) -> Duration {
self.stream.timeout()
}
async fn send(&mut self, packet: P) -> io::Result<()> {
let mut bytes = packet.into_bytes();
bytes.encrypt(&mut self.send_key);
let slice = bytes.as_slice();
self.stream.write_all(slice).await?;
self.stream.flush().await?;
Ok(())
}
async fn receive(&mut self) -> Result<P, PacketReceiverError<P::Header>> {
let recv_key = &mut self.recv_key;
self.builder
.read_header(&mut self.stream, |bytes| {
bytes.decrypt_header(recv_key).map_err(|e| e.into())
})
.await?;
self.builder
.read_body(&mut self.stream, |bytes| {
bytes.decrypt_body(recv_key).map_err(|e| e.into())
})
.await
}
async fn shutdown(&mut self) -> io::Result<()> {
self.stream.shutdown().await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{packet::test::TestPacket, server::Request};
use crypto::signature::Keypair;
use tokio::net::{TcpListener, TcpStream};
async fn tcp_streams() -> (TcpStream, TcpStream) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let connect = TcpStream::connect(addr);
let accept = listener.accept();
let (connect, accept) = tokio::join!(connect, accept);
(connect.unwrap(), accept.unwrap().0)
}
#[tokio::test]
async fn test_encrypted_stream() {
let timeout = Duration::from_secs(1);
let key = Keypair::new();
let public_key = key.public().clone();
let (alice, bob) = tcp_streams().await;
let alice: Client<TestPacket<_>> = client(
alice,
ClientConfig {
timeout,
body_limit: 200,
},
None,
public_key,
);
let bob_task = tokio::spawn(async move {
let mut bob: Server<TestPacket<_>> = server(
bob,
ServerConfig {
timeout,
body_limit: 200,
},
key,
);
let req = bob.receive().await.unwrap();
match req {
Request::Request(req, resp) => {
assert_eq!(req.num1, 1);
assert_eq!(req.num2, 2);
let res = TestPacket::new(3, 4);
resp.send(res).unwrap();
}
_ => panic!("expected request"),
};
let req = bob.receive().await.unwrap();
match req {
Request::RequestReceiver(req, stream) => {
assert_eq!(req.num1, 5);
assert_eq!(req.num2, 6);
let res = TestPacket::new(7, 8);
stream.send(res).await.unwrap();
let res = TestPacket::new(9, 10);
stream.send(res).await.unwrap();
}
_ => panic!("expected stream"),
};
let req = bob.receive().await.unwrap();
match req {
Request::RequestSender(req, mut stream) => {
assert_eq!(req.num1, 11);
assert_eq!(req.num2, 12);
let res = stream.receive().await.unwrap();
assert_eq!(res.num1, 13);
assert_eq!(res.num2, 14);
let res = stream.receive().await.unwrap();
assert_eq!(res.num1, 15);
assert_eq!(res.num2, 16);
}
_ => panic!("expected stream"),
};
bob.wait().await.unwrap();
});
let req = TestPacket::new(1, 2);
let res = alice.request(req).await.unwrap();
assert_eq!(res.num1, 3);
assert_eq!(res.num2, 4);
let req = TestPacket::new(5, 6);
let mut stream = alice.request_receiver(req).await.unwrap();
let res = stream.receive().await.unwrap();
assert_eq!(res.num1, 7);
assert_eq!(res.num2, 8);
let res = stream.receive().await.unwrap();
assert_eq!(res.num1, 9);
assert_eq!(res.num2, 10);
drop(stream);
let req = TestPacket::new(11, 12);
let stream = alice.request_sender(req).await.unwrap();
let req = TestPacket::new(13, 14);
stream.send(req).await.unwrap();
let req = TestPacket::new(15, 16);
stream.send(req).await.unwrap();
drop(stream);
alice.close().await.unwrap();
bob_task.await.unwrap();
}
}