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 {
pub fn from_addr(addr: SocketAddr) -> BtpConfig {
BtpConfig {
ver: super::VERSION,
addr,
}
}
}
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(())
}
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,
}
}
}