lafere 0.2.0-pre.6

A more or less simple communication protocol library.
Documentation
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)
}

/// inner manages a stream
struct EncryptedPacketStream<S, P>
where
	S: ByteStream,
	P: Packet<EncryptedBytes>,
{
	stream: TimeoutReader<S>,
	send_key: Key,
	recv_key: Key,
	// buffer to receive a message
	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(())
	}

	/// this function is abort safe
	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's receive a request message
			let req = bob.receive().await.unwrap();
			match req {
				Request::Request(req, resp) => {
					assert_eq!(req.num1, 1);
					assert_eq!(req.num2, 2);

					// send response
					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);

					// send response
					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);

					// send response
					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's make a request
		let req = TestPacket::new(1, 2);
		let res = alice.request(req).await.unwrap();
		assert_eq!(res.num1, 3);
		assert_eq!(res.num2, 4);

		// let's create a stream to listen
		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);

		// now request a stream.sender
		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();

		// wait until bob's task finishes
		bob_task.await.unwrap();
	}
}