use super::socket::{BtpConfig, BtpSocket};
use std::future::Future;
use std::net::SocketAddr;
use tokio::net::{TcpListener, TcpStream};
#[cfg(feature = "tls")]
use std::sync::Arc;
type Error = Box<dyn std::error::Error + Send + Sync>;
pub trait AsListener {
fn accept(&self) -> impl Future<Output = Result<(TcpStream, SocketAddr), std::io::Error>>;
}
#[cfg(feature = "tls")]
pub struct BtpListenerTls<Listener: AsListener> {
pub listener: Listener,
pub btp_conf: BtpConfig,
pub tls_conf: Arc<rustls::ServerConfig>,
}
pub struct BtpListener<Listener: AsListener> {
pub listener: Listener,
pub btp_conf: BtpConfig,
}
impl AsListener for TcpListener {
fn accept(&self) -> impl Future<Output = Result<(TcpStream, SocketAddr), std::io::Error>> {
return self.accept();
}
}
mod _impl {
use tokio::{
io::{ReadHalf, WriteHalf},
spawn,
};
use super::*;
impl<L: AsListener> BtpListener<L> {
pub async fn from(listener: L, btp_conf: BtpConfig) -> Self {
BtpListener { listener, btp_conf }
}
pub async fn accept(
&self,
) -> Result<BtpSocket<ReadHalf<TcpStream>, WriteHalf<TcpStream>, std::net::SocketAddr>, Error>
{
let (socket, peer) = self.listener.accept().await?;
let me = BtpConfig {
ver: self.btp_conf.ver,
addr: self.btp_conf.addr,
};
let mut socket = BtpSocket::from(socket, me, peer);
let mut inner = socket.inner.clone();
let sender = socket.queue.as_ref().unwrap().0.clone();
spawn(async move {
loop {
match inner.read().await {
Ok(packet) => {
sender.send(packet).await.unwrap();
}
Err(e) => {
eprintln!("Couldn't send the packet: {}", e);
}
};
}
});
Ok(socket)
}
}
}
#[cfg(feature = "tls")]
mod _impl_tls {
use crate::{
VERSION,
message::{BtpPackage, StatusCode},
};
use super::*;
use rustls::ServerConfig;
use std::{
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
time::Duration,
};
use tokio::{
io::{ReadHalf, WriteHalf},
spawn,
time::sleep,
};
use tokio_rustls::{TlsAcceptor, server::TlsStream};
impl<L: AsListener> BtpListenerTls<L> {
pub async fn from(listener: L, btp_conf: BtpConfig, tls_conf: Arc<ServerConfig>) -> Self {
BtpListenerTls {
listener,
btp_conf,
tls_conf,
}
}
pub async fn accept(
&self,
) -> Result<
BtpSocket<
ReadHalf<TlsStream<TcpStream>>,
WriteHalf<TlsStream<TcpStream>>,
std::net::SocketAddr,
>,
Error,
> {
let acceptor = TlsAcceptor::from(self.tls_conf.clone());
let (socket, peer) = self.listener.accept().await?;
let tls_stream = acceptor.accept(socket).await?;
let conf = BtpConfig {
ver: VERSION,
addr: peer,
};
let mut socket = BtpSocket::from(tls_stream, conf, peer);
let clone = socket.copy_it();
let mut clone_t = socket.copy_it();
let sender = socket.queue.as_ref().unwrap().0.clone();
let mut inner = socket.inner.clone();
let close = Arc::new(AtomicBool::new(true));
let close_r = Arc::clone(&close);
let close_t = Arc::clone(&close);
spawn(async move {
sleep(Duration::from_secs(25)).await;
if close_r.load(Ordering::SeqCst) {
clone.shutdown().await.unwrap();
} else {
close_r.store(true, Ordering::SeqCst);
}
});
spawn(async move {
loop {
match inner.read().await {
Ok(packet) => {
if packet.header.get_sc() == StatusCode::OkieDokie
&& packet.body1 == "okie"
{
close_t.store(false, Ordering::SeqCst);
clone_t
.write(BtpPackage::dokie(clone_t.btp_conf.ver))
.await
.unwrap();
} else {
sender.send(packet).await.unwrap();
}
}
Err(e) => {
eprintln!("Couldn't send the packet: {}", e);
}
};
}
});
Ok(socket)
}
}
}