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>;
pub struct BtpConfig {
pub ver: [u8; 3],
pub addr: SocketAddr,
}
pub struct BtpSocketInner<
S: AsyncRead + AsyncWrite + Unpin + Send + Sync,
P: ToSocketAddrs + Send + Sync,
> {
pub socket: S,
pub btp_conf: BtpConfig,
pub peer: P,
}
#[derive(Clone)]
pub struct BtpSocket<
S: AsyncRead + AsyncWrite + Unpin + Send + Sync,
P: ToSocketAddrs + Send + Sync,
> {
inner: Arc<Mutex<BtpSocketInner<S, P>>>,
}
impl<S: AsyncRead + AsyncWrite + Unpin + Send + Sync, P: ToSocketAddrs + Send + Sync>
BtpSocketInner<S, P>
{
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 + 'static,
> BtpSocket<S, P>
{
pub async fn write_message(&mut self, msg: BtpMessage) -> Result<(), Error> {
let mut inner = self.inner.lock().await;
let mut writer = BufWriter::new(&mut inner.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 inner = self.inner.lock().await;
let mut writer = BufWriter::new(&mut inner.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 inner = self.inner.lock().await;
let mut reader = BufReader::new(&mut inner.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
}
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) + Send + Sync + 'static,
F2: Fn(BtpFile) + Send + Sync + 'static,
{
let inner = Arc::clone(&self.inner);
let func_msg = Arc::new(func_msg);
let func_file = Arc::new(func_file);
spawn(async move {
loop {
match inner.lock().await.read().await {
Ok((None, Some(msg))) => {
let f = Arc::clone(&func_msg);
std::thread::spawn(move || f(msg));
}
Ok((Some(file), None)) => {
let f = Arc::clone(&func_file);
std::thread::spawn(move || f(file));
}
Err(e) => {
eprintln!("Couldn't send the message: {}", e);
}
_ => {}
};
}
});
Ok(())
}
}