mod _impl {
type Error = Box<dyn std::error::Error + Send + Sync>;
use crate::{
VERSION,
message::{BtpPackage, StatusCode},
socket::{BtpConfig, BtpSocket},
};
use std::{
net::SocketAddr,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
time::Duration,
};
use tokio::io::{ReadHalf, WriteHalf};
use tokio::{
net::{TcpStream, lookup_host},
spawn,
time::sleep,
};
#[cfg(feature = "tls")]
use rustls::pki_types::ServerName;
#[cfg(feature = "tls")]
use tokio_rustls::{TlsConnector, client::TlsStream};
impl BtpSocket<ReadHalf<TcpStream>, WriteHalf<TcpStream>, SocketAddr> {
pub async fn connect(
conf: BtpConfig,
) -> Result<BtpSocket<ReadHalf<TcpStream>, WriteHalf<TcpStream>, SocketAddr>, Error>
{
let sock_addr = lookup_host(&conf.addr)
.await?
.next()
.ok_or("Couldn't find the IP address of the host")?;
let close = Arc::new(AtomicBool::new(false));
let close_r = Arc::clone(&close);
let close_w = Arc::clone(&close);
let sock = TcpStream::connect(&conf.addr).await.unwrap();
let mut btp_socket = BtpSocket::from(sock, conf, sock_addr);
let mut clone = btp_socket.copy_it();
let sender = btp_socket.queue.as_ref().unwrap().0.clone();
let mut inner = btp_socket.inner.clone();
let okier: tokio::task::JoinHandle<Result<(), Error>> = spawn(async move {
sleep(Duration::from_secs(1)).await;
loop {
if close_r.load(Ordering::SeqCst) {
clone.shutdown().await.unwrap();
break;
}
close_r.store(true, Ordering::SeqCst);
clone.write(BtpPackage::okie(VERSION)).await?;
sleep(Duration::from_secs(20)).await;
}
Ok(())
});
let reader: tokio::task::JoinHandle<Result<(), Error>> = spawn(async move {
sleep(Duration::from_secs(1)).await;
loop {
match inner.read().await {
Ok(packet) => {
if packet.header.get_sc() == StatusCode::OkieDokie
&& packet.body1 == "dokie"
{
close_w.store(true, Ordering::SeqCst);
} else {
sender.send(packet).await.unwrap();
}
}
Err(e) => {
eprintln!("Couldn't send the packet: {}", e);
break;
}
};
}
Ok(())
});
let clone = btp_socket.copy_it();
let mut d_okier = clone.inner.d_okier.lock().await;
*d_okier = Some((okier, reader));
Ok(btp_socket)
}
pub async fn connect_without_okie_dokie(
conf: BtpConfig,
) -> Result<BtpSocket<ReadHalf<TcpStream>, WriteHalf<TcpStream>, SocketAddr>, Error>
{
let sock_addr = lookup_host(&conf.addr)
.await?
.next()
.ok_or("Couldn't find the IP address of the host")?;
let sock = TcpStream::connect(&conf.addr).await.unwrap();
Ok(BtpSocket::from(sock, conf, sock_addr))
}
#[cfg(feature = "tls")]
pub async fn connect_tls(
connector: TlsConnector,
conf: BtpConfig,
) -> Result<
BtpSocket<ReadHalf<TlsStream<TcpStream>>, WriteHalf<TlsStream<TcpStream>>, SocketAddr>,
Error,
> {
let close = Arc::new(AtomicBool::new(false));
let close_r = Arc::clone(&close);
let close_w = Arc::clone(&close);
let sock_addr = lookup_host(&conf.addr)
.await?
.next()
.ok_or("Couldn't find the IP address of the host")?;
let server_name = ServerName::IpAddress(sock_addr.ip().into());
let sock = TcpStream::connect(&conf.addr).await.unwrap();
let sock = connector.connect(server_name, sock).await?;
let mut btp_socket = BtpSocket::from(sock, conf, sock_addr);
let mut clone = btp_socket.copy_it();
let sender = btp_socket.queue.as_ref().unwrap().0.clone();
let mut inner = btp_socket.inner.clone();
let okier: tokio::task::JoinHandle<Result<(), Error>> = spawn(async move {
sleep(Duration::from_secs(1)).await;
loop {
if close_r.load(Ordering::SeqCst) {
clone.shutdown().await.unwrap();
break;
}
close_r.store(true, Ordering::SeqCst);
clone.write(BtpPackage::okie(VERSION)).await?;
sleep(Duration::from_secs(20)).await;
}
Ok(())
});
let reader: tokio::task::JoinHandle<Result<(), Error>> = spawn(async move {
sleep(Duration::from_secs(1)).await;
loop {
match inner.read().await {
Ok(packet) => {
if packet.header.get_sc() == StatusCode::OkieDokie
&& packet.body1 == "dokie"
{
close_w.store(true, Ordering::SeqCst);
} else {
sender.send(packet).await.unwrap();
}
}
Err(e) => {
eprintln!("Couldn't send the packet: {}", e);
break;
}
};
}
Ok(())
});
let clone = btp_socket.copy_it();
let mut d_okier = clone.inner.d_okier.lock().await;
*d_okier = Some((okier, reader));
Ok(btp_socket)
}
}
}