use crate::message::{BtpFile, BtpHeader, BtpMessage};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader, BufWriter};
use tokio::net::ToSocketAddrs;
use tokio::spawn;
use tokio::sync::Mutex;
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,
}
}
}
struct BtpSocketInner<S: AsyncRead + AsyncWrite + Unpin + Send + Sync> {
pub socket: S,
}
#[derive(Clone)]
pub struct BtpSocket<
S: AsyncRead + AsyncWrite + Unpin + Send + Sync,
P: ToSocketAddrs + Send + Sync,
> {
inner: Arc<Mutex<BtpSocketInner<S>>>,
pub btp_conf: BtpConfig,
pub peer: P,
}
impl<S: AsyncRead + AsyncWrite + Unpin + Send + Sync> BtpSocketInner<S> {
pub async fn write_message(&mut self, msg: BtpMessage) -> Result<(), Error> {
let mut writer = BufWriter::new(&mut self.socket);
writer.write_all(msg.as_vec().as_slice()).await?;
writer.flush().await?;
Ok(())
}
pub async fn write_file(&mut self, msg: BtpFile) -> Result<(), Error> {
let mut writer = BufWriter::new(&mut self.socket);
writer.write_all(msg.as_vec().as_slice()).await?;
writer.flush().await?;
Ok(())
}
pub async fn read(&mut self) -> Result<(Option<BtpFile>, Option<BtpMessage>), Error> {
let mut reader = BufReader::new(&mut self.socket);
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?;
reader.flush().await?;
res = Ok((
None,
Some(BtpMessage::from_header(header, String::from_utf8(buf)?)),
));
}
_ => {
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.flush().await?;
reader.read_exact(&mut file_name).await?;
reader.flush().await?;
res = Ok((
Some(BtpFile::from_header(
header,
String::from_utf8(file_name)?,
String::from_utf8(file)?,
)),
None,
));
}
}
res
}
}
impl<
S: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static,
P: ToSocketAddrs + Send + Sync + Clone + Copy + 'static,
> BtpSocket<S, P>
{
pub async fn write_message(&mut self, msg: BtpMessage) -> Result<(), Error> {
let mut inner = self.inner.lock().await;
inner.write_message(msg).await
}
pub async fn write_file(&mut self, msg: BtpFile) -> Result<(), Error> {
let mut inner = self.inner.lock().await;
inner.write_file(msg).await
}
pub async fn read(&mut self) -> Result<(Option<BtpFile>, Option<BtpMessage>), Error> {
let mut inner = self.inner.lock().await;
inner.read().await
}
pub fn from(socket: S, btp_conf: BtpConfig, peer: P) -> Self {
BtpSocket {
inner: Arc::new(Mutex::new(BtpSocketInner { socket })),
btp_conf,
peer,
}
}
pub fn attach_handler<F1, F2>(&self, func_msg: F1, func_file: F2) -> Result<(), Error>
where
F1: Fn(BtpMessage, BtpSocket<S, P>) + Send + Sync + 'static,
F2: Fn(BtpFile, BtpSocket<S, P>) + Send + Sync + 'static,
{
let inner = Arc::clone(&self.inner);
let func_msg = Arc::new(func_msg);
let func_file = Arc::new(func_file);
let peer = self.peer.clone();
let btp_conf = self.btp_conf.clone();
spawn(async move {
loop {
match inner.lock().await.read().await {
Ok((None, Some(msg))) => {
let f = Arc::clone(&func_msg);
let i = Arc::clone(&inner);
let soc = BtpSocket {
inner: i,
btp_conf,
peer,
};
std::thread::spawn(move || f(msg, soc));
}
Ok((Some(file), None)) => {
let f = Arc::clone(&func_file);
let i = Arc::clone(&inner);
let soc = BtpSocket {
inner: i,
btp_conf,
peer,
};
std::thread::spawn(move || f(file, soc));
}
Err(e) => {
eprintln!("Couldn't send the message: {}", e);
}
_ => {}
};
}
});
Ok(())
}
}