btp 0.1.4

A rust library of, a lightweight protocol, Blog Transfer Protocol.
Documentation
use crate::message::{BtpHeader, BtpPackage};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadHalf, WriteHalf, split};
use tokio::net::ToSocketAddrs;
use tokio::sync::Mutex;
use tokio::sync::mpsc::{Receiver, Sender, channel};
use tokio::task::JoinHandle;

type Error = Box<dyn std::error::Error + Send + Sync>;

#[derive(Clone, Copy)]
pub struct BtpConfig {
    pub ver: [u8; 3],
    pub addr: SocketAddr,
}

impl BtpConfig {
    /// Takes the address and returns [`BtpConfig`]
    pub fn from_addr(addr: SocketAddr) -> BtpConfig {
        BtpConfig {
            ver: super::VERSION,
            addr,
        }
    }
}

/// Since we use this function inside the async functions, we need to send, sync it between threads,
/// so we use this [`BtpSocketInner`] wrapped in an Arc<Mutex<>> to use between threads safely without any additional problem.
/// We don't use [`BtpSocketInner`] inside the functions, it's not public because it's embedded into [`BtpSocket`]
pub(crate) struct BtpSocketInner<
    R: AsyncRead + Unpin + Send + Sync,
    W: AsyncWrite + Unpin + Send + Sync,
> {
    pub reader: Arc<Mutex<R>>,
    pub writer: Arc<Mutex<W>>,
    pub d_okier: Arc<Mutex<Option<(JoinHandle<Result<(), Error>>, JoinHandle<Result<(), Error>>)>>>,
}

pub struct BtpSocket<
    R: AsyncRead + Unpin + Send + Sync,
    W: AsyncWrite + Unpin + Send + Sync,
    P: ToSocketAddrs + Send + Sync,
> {
    pub(crate) inner: BtpSocketInner<R, W>,
    pub btp_conf: BtpConfig,
    pub peer: P,
    pub(crate) queue: Option<(Sender<BtpPackage>, Receiver<BtpPackage>)>,
}

impl<R: AsyncRead + Unpin + Send + Sync, W: AsyncWrite + Unpin + Send + Sync> BtpSocketInner<R, W> {
    pub async fn write_message(&mut self, msg: BtpPackage) -> Result<(), Error> {
        let writerc = Arc::clone(&self.writer);

        let mut writer = writerc.lock().await;
        writer.write_all(msg.as_vec().as_slice()).await?;
        writer.flush().await?;

        Ok(())
    }

    /// Returns a Result and a [`BtpFile`] or [`BtpMessage`], with checking the file_name_len.
    /// If it's empty, it understands that coming message is BtpMessage.
    /// If it has a file_name_len, then it is a BtpFile.
    pub async fn read(&mut self) -> Result<BtpPackage, Error> {
        let readerc = Arc::clone(&self.reader);

        let mut reader = readerc.lock().await;
        let mut header = [0u8; 10];
        reader.read_exact(&mut header).await?;
        let res;
        match &header[8..] {
            [0, 0] => {
                let header = BtpHeader::from_raw(header)?;
                let file_len = header.get_len();
                let mut buf = vec![0u8; file_len];
                reader.read_exact(&mut buf).await?;
                res = Ok(BtpPackage::from_header(
                    header,
                    String::from_utf8(buf)?,
                    None,
                ));
            }
            _ => {
                let header = BtpHeader::from_raw(header)?;
                let file_len = header.get_len();
                let file_name_len = header.get_file_name_len();
                let mut file = vec![0u8; file_len];
                let mut file_name = vec![0u8; file_name_len];
                reader.read_exact(&mut file).await?;
                reader.read_exact(&mut file_name).await?;
                res = Ok(BtpPackage::from_header(
                    header,
                    String::from_utf8(file)?,
                    Some(String::from_utf8(file_name)?),
                ));
            }
        }
        res
    }

    pub fn clone(&mut self) -> Self {
        let reader = Arc::clone(&self.reader);
        let writer = Arc::clone(&self.writer);
        let d_okier = Arc::clone(&self.d_okier);
        BtpSocketInner {
            reader,
            writer,
            d_okier,
        }
    }
}

impl<
    R: AsyncRead + Unpin + Send + Sync,
    W: AsyncWrite + Unpin + Send + Sync,
    P: ToSocketAddrs + Send + Sync + Clone + Copy + 'static,
> BtpSocket<R, W, P>
{
    pub(super) fn copy_it(&mut self) -> Self {
        let peer = self.peer.clone();
        let btp_conf = self.btp_conf.clone();
        let inner = self.inner.clone();

        BtpSocket {
            inner,
            btp_conf,
            peer,
            queue: None,
        }
    }

    pub async fn shutdown(self) -> Result<(), Error> {
        if let Some((read_task, write_task)) = self.inner.d_okier.lock().await.take() {
            read_task.abort();
            write_task.abort();
        }

        let mut writer = self.inner.writer.lock().await;
        writer.shutdown().await?;

        Ok(())
    }

    pub async fn write(&mut self, msg: BtpPackage) -> Result<(), Error> {
        self.inner.write_message(msg).await
    }

    pub async fn read(&mut self) -> Result<Option<BtpPackage>, Error> {
        match self.queue.as_mut() {
            Some((_, r)) => Ok(r.recv().await),
            None => panic!("You can't read messages from a copy."),
        }
    }
}

impl<S, P> BtpSocket<ReadHalf<S>, WriteHalf<S>, P>
where
    S: AsyncRead + AsyncWrite + Unpin + Send + Sync,
    P: ToSocketAddrs + Send + Sync + Clone + Copy + 'static,
{
    pub fn from(socket: S, btp_conf: BtpConfig, peer: P) -> Self {
        let queue = Some(channel::<BtpPackage>(128));
        let (readerc, writerc) = split(socket);
        let (reader, writer) = (Arc::new(Mutex::new(readerc)), Arc::new(Mutex::new(writerc)));

        BtpSocket {
            inner: BtpSocketInner {
                reader,
                writer,
                d_okier: Arc::new(Mutex::new(None)),
            },
            btp_conf,
            peer,
            queue,
        }
    }
}